mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-27 08:00:17 +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
+1
-1
@@ -37,7 +37,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| Session management/routing | ✅ | ✅ | SessionManager exists |
|
||||
| Configuration hot-reload | ✅ | ❌ | |
|
||||
| Network modes (loopback/LAN/remote) | ✅ | 🚧 | HTTP only |
|
||||
| OpenAI-compatible HTTP API | ✅ | ✅ | /v1/chat/completions |
|
||||
| OpenAI-compatible HTTP API | ✅ | ✅ | /v1/chat/completions, per-request `model` override |
|
||||
| Canvas hosting | ✅ | ❌ | Agent-driven UI |
|
||||
| Gateway lock (PID-based) | ✅ | ❌ | |
|
||||
| launchd/systemd integration | ✅ | ❌ | |
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -24,7 +24,21 @@ const AUTH_TOKEN: &str = "test-openai-token";
|
||||
// Mock LLM provider
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
struct MockLlmProvider;
|
||||
#[derive(Default)]
|
||||
struct MockLlmState {
|
||||
completion_models: tokio::sync::Mutex<Vec<Option<String>>>,
|
||||
tool_completion_models: tokio::sync::Mutex<Vec<Option<String>>>,
|
||||
}
|
||||
|
||||
struct MockLlmProvider {
|
||||
state: Arc<MockLlmState>,
|
||||
}
|
||||
|
||||
impl MockLlmProvider {
|
||||
fn new(state: Arc<MockLlmState>) -> Self {
|
||||
Self { state }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for MockLlmProvider {
|
||||
@@ -37,6 +51,12 @@ impl LlmProvider for MockLlmProvider {
|
||||
}
|
||||
|
||||
async fn complete(&self, req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
|
||||
self.state
|
||||
.completion_models
|
||||
.lock()
|
||||
.await
|
||||
.push(req.model.clone());
|
||||
|
||||
// Echo the last user message back
|
||||
let user_msg = req
|
||||
.messages
|
||||
@@ -59,6 +79,12 @@ impl LlmProvider for MockLlmProvider {
|
||||
&self,
|
||||
req: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, LlmError> {
|
||||
self.state
|
||||
.tool_completion_models
|
||||
.lock()
|
||||
.await
|
||||
.push(req.model.clone());
|
||||
|
||||
// If tools are provided, return a tool call
|
||||
if let Some(tool) = req.tools.first() {
|
||||
Ok(ToolCompletionResponse {
|
||||
@@ -93,11 +119,71 @@ impl LlmProvider for MockLlmProvider {
|
||||
}
|
||||
}
|
||||
|
||||
struct FixedModelProvider {
|
||||
model: &'static str,
|
||||
}
|
||||
|
||||
impl FixedModelProvider {
|
||||
fn new(model: &'static str) -> Self {
|
||||
Self { model }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for FixedModelProvider {
|
||||
fn model_name(&self) -> &str {
|
||||
self.model
|
||||
}
|
||||
|
||||
fn cost_per_token(&self) -> (Decimal, Decimal) {
|
||||
(Decimal::ZERO, Decimal::ZERO)
|
||||
}
|
||||
|
||||
async fn complete(&self, _req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
|
||||
Ok(CompletionResponse {
|
||||
content: "fixed response".to_string(),
|
||||
input_tokens: 10,
|
||||
output_tokens: 5,
|
||||
finish_reason: FinishReason::Stop,
|
||||
response_id: None,
|
||||
})
|
||||
}
|
||||
|
||||
async fn complete_with_tools(
|
||||
&self,
|
||||
_req: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, LlmError> {
|
||||
Ok(ToolCompletionResponse {
|
||||
content: Some("fixed response".to_string()),
|
||||
tool_calls: vec![],
|
||||
input_tokens: 10,
|
||||
output_tokens: 5,
|
||||
finish_reason: FinishReason::Stop,
|
||||
response_id: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn effective_model_name(&self, _requested_model: Option<&str>) -> String {
|
||||
self.model.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Test helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
async fn start_test_server() -> (SocketAddr, Arc<GatewayState>) {
|
||||
async fn start_test_server() -> (SocketAddr, Arc<GatewayState>, Arc<MockLlmState>) {
|
||||
let mock_state = Arc::new(MockLlmState::default());
|
||||
|
||||
let llm_provider: Arc<dyn LlmProvider> = Arc::new(MockLlmProvider::new(mock_state.clone()));
|
||||
let (bound_addr, state) = start_test_server_with_provider(llm_provider).await;
|
||||
|
||||
(bound_addr, state, mock_state)
|
||||
}
|
||||
|
||||
async fn start_test_server_with_provider(
|
||||
llm_provider: Arc<dyn LlmProvider>,
|
||||
) -> (SocketAddr, Arc<GatewayState>) {
|
||||
let state = Arc::new(GatewayState {
|
||||
msg_tx: tokio::sync::RwLock::new(None),
|
||||
sse: SseManager::new(),
|
||||
@@ -112,7 +198,7 @@ async fn start_test_server() -> (SocketAddr, Arc<GatewayState>) {
|
||||
user_id: "test-user".to_string(),
|
||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
||||
llm_provider: Some(Arc::new(MockLlmProvider)),
|
||||
llm_provider: Some(llm_provider),
|
||||
skill_registry: None,
|
||||
skill_catalog: None,
|
||||
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
|
||||
@@ -139,7 +225,7 @@ fn client() -> reqwest::Client {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_basic() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
let (addr, _state, mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
@@ -173,11 +259,14 @@ async fn test_chat_completions_basic() {
|
||||
assert_eq!(body["usage"]["prompt_tokens"], 10);
|
||||
assert_eq!(body["usage"]["completion_tokens"], 5);
|
||||
assert_eq!(body["usage"]["total_tokens"], 15);
|
||||
|
||||
let models = mock_state.completion_models.lock().await;
|
||||
assert_eq!(*models, vec![Some("mock-model-v1".to_string())]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_with_system_message() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
let (addr, _state, _mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
@@ -204,7 +293,7 @@ async fn test_chat_completions_with_system_message() {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_with_tools() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
let (addr, _state, mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
@@ -243,11 +332,14 @@ async fn test_chat_completions_with_tools() {
|
||||
assert_eq!(tool_calls[0]["id"], "call_mock_001");
|
||||
assert_eq!(tool_calls[0]["type"], "function");
|
||||
assert_eq!(tool_calls[0]["function"]["name"], "get_weather");
|
||||
|
||||
let models = mock_state.tool_completion_models.lock().await;
|
||||
assert_eq!(*models, vec![Some("mock-model-v1".to_string())]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_streaming() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
let (addr, _state, mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
@@ -316,11 +408,14 @@ async fn test_chat_completions_streaming() {
|
||||
"Expected reassembled content to contain 'Stream test', got: '{}'",
|
||||
full_content
|
||||
);
|
||||
|
||||
let models = mock_state.completion_models.lock().await;
|
||||
assert_eq!(*models, vec![Some("mock-model-v1".to_string())]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_empty_messages() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
let (addr, _state, _mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
@@ -340,8 +435,8 @@ async fn test_chat_completions_empty_messages() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_model_mismatch() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
async fn test_chat_completions_model_override() {
|
||||
let (addr, _state, mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
@@ -355,20 +450,173 @@ async fn test_chat_completions_model_mismatch() {
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(resp.status(), 404);
|
||||
assert_eq!(resp.status(), 200);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
assert_eq!(body["model"], "gpt-4");
|
||||
|
||||
let models = mock_state.completion_models.lock().await;
|
||||
assert_eq!(*models, vec![Some("gpt-4".to_string())]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_uses_effective_model_when_override_ignored() {
|
||||
let provider: Arc<dyn LlmProvider> = Arc::new(FixedModelProvider::new("configured-model"));
|
||||
let (addr, _state) = start_test_server_with_provider(provider).await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
.post(&url)
|
||||
.bearer_auth(AUTH_TOKEN)
|
||||
.json(&serde_json::json!({
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hi"}]
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(resp.status(), 200);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
assert_eq!(body["model"], "configured-model");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_streaming_uses_effective_model_when_override_ignored() {
|
||||
let provider: Arc<dyn LlmProvider> = Arc::new(FixedModelProvider::new("configured-model"));
|
||||
let (addr, _state) = start_test_server_with_provider(provider).await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
.post(&url)
|
||||
.bearer_auth(AUTH_TOKEN)
|
||||
.json(&serde_json::json!({
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hi"}],
|
||||
"stream": true
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(resp.status(), 200);
|
||||
let text = resp.text().await.unwrap();
|
||||
assert!(
|
||||
text.contains("\"model\":\"configured-model\""),
|
||||
"Expected streaming chunks to report configured model, got: {}",
|
||||
text
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_model_too_long() {
|
||||
let (addr, _state, mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
.post(&url)
|
||||
.bearer_auth(AUTH_TOKEN)
|
||||
.json(&serde_json::json!({
|
||||
"model": "m".repeat(300),
|
||||
"messages": [{"role": "user", "content": "Hi"}]
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(resp.status(), 400);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
assert_eq!(body["error"]["code"], "model_not_found");
|
||||
assert!(
|
||||
body["error"]["message"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("mock-model-v1")
|
||||
.unwrap_or("")
|
||||
.contains("model"),
|
||||
"Expected model validation error, got: {}",
|
||||
body
|
||||
);
|
||||
|
||||
// Validation should fail before provider invocation.
|
||||
let models = mock_state.completion_models.lock().await;
|
||||
assert!(
|
||||
models.is_empty(),
|
||||
"provider should not be called: {:?}",
|
||||
*models
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_model_with_control_chars() {
|
||||
let (addr, _state, mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
.post(&url)
|
||||
.bearer_auth(AUTH_TOKEN)
|
||||
.json(&serde_json::json!({
|
||||
"model": "gpt-4\noops",
|
||||
"messages": [{"role": "user", "content": "Hi"}]
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(resp.status(), 400);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
assert!(
|
||||
body["error"]["message"]
|
||||
.as_str()
|
||||
.unwrap_or("")
|
||||
.contains("control"),
|
||||
"Expected model validation error, got: {}",
|
||||
body
|
||||
);
|
||||
|
||||
// Validation should fail before provider invocation.
|
||||
let models = mock_state.completion_models.lock().await;
|
||||
assert!(
|
||||
models.is_empty(),
|
||||
"provider should not be called: {:?}",
|
||||
*models
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_model_with_surrounding_whitespace() {
|
||||
let (addr, _state, mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
.post(&url)
|
||||
.bearer_auth(AUTH_TOKEN)
|
||||
.json(&serde_json::json!({
|
||||
"model": " gpt-4 ",
|
||||
"messages": [{"role": "user", "content": "Hi"}]
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(resp.status(), 400);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
assert!(
|
||||
body["error"]["message"]
|
||||
.as_str()
|
||||
.unwrap_or("")
|
||||
.contains("leading or trailing whitespace"),
|
||||
"Expected model validation error, got: {}",
|
||||
body
|
||||
);
|
||||
|
||||
let models = mock_state.completion_models.lock().await;
|
||||
assert!(
|
||||
models.is_empty(),
|
||||
"provider should not be called: {:?}",
|
||||
*models
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_no_auth() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
let (addr, _state, _mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
let resp = client()
|
||||
@@ -387,7 +635,7 @@ async fn test_chat_completions_no_auth() {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_models_endpoint() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
let (addr, _state, _mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/models", addr);
|
||||
|
||||
let resp = client()
|
||||
@@ -410,7 +658,7 @@ async fn test_models_endpoint() {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_models_no_auth() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
let (addr, _state, _mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/models", addr);
|
||||
|
||||
let resp = client().get(&url).send().await.unwrap();
|
||||
@@ -462,7 +710,7 @@ async fn test_no_llm_provider_returns_503() {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_chat_completions_body_too_large() {
|
||||
let (addr, _state) = start_test_server().await;
|
||||
let (addr, _state, _mock_state) = start_test_server().await;
|
||||
let url = format!("http://{}/v1/chat/completions", addr);
|
||||
|
||||
// Build a payload over 1 MB (the gateway's DefaultBodyLimit)
|
||||
|
||||
Reference in New Issue
Block a user