Use X-SMG-Routing-Key header instead of json body and add tests (#15826)

This commit is contained in:
fzyzcjy
2025-12-25 09:58:14 -08:00
committed by GitHub
parent a355911909
commit 6ef543f99d
22 changed files with 106 additions and 107 deletions
@@ -67,7 +67,6 @@ fn default_generate_request() -> GenerateRequest {
return_bytes: false, return_bytes: false,
return_entropy: false, return_entropy: false,
rid: None, rid: None,
routing_id: None,
} }
} }
@@ -123,7 +122,6 @@ fn default_completion_request() -> CompletionRequest {
return_hidden_states: false, return_hidden_states: false,
sampling_seed: None, sampling_seed: None,
other: serde_json::Map::new(), other: serde_json::Map::new(),
routing_id: None,
} }
} }
@@ -0,0 +1,52 @@
import collections
import pytest
import requests
ROUTING_KEY_HEADER = "X-SMG-Routing-Key"
@pytest.mark.integration
def test_manual_routing_with_header(mock_workers, router_manager):
"""With X-SMG-Routing-Key header: sticky routing + distribution across workers."""
_, urls, _ = mock_workers(n=2)
rh = router_manager.start_router(worker_urls=urls, policy="manual")
# Send requests: 5 keys × 4 requests each
results = collections.defaultdict(set)
with requests.Session() as s:
for key_id in range(5):
for _ in range(4):
worker = send_completion(s, rh.url, f"user-{key_id}")
results[f"user-{key_id}"].add(worker)
# Verify sticky: each key should route to exactly one worker
for key, workers in results.items():
assert len(workers) == 1, f"Key {key} routed to multiple workers: {workers}"
# Verify distribution: different keys should use multiple workers
all_workers = {list(w)[0] for w in results.values()}
assert len(all_workers) > 1, f"Should distribute across workers: {results}"
@pytest.mark.integration
def test_manual_routing_without_header(mock_workers, router_manager):
"""Without X-SMG-Routing-Key header: random fallback distribution."""
_, urls, _ = mock_workers(n=2)
rh = router_manager.start_router(worker_urls=urls, policy="manual")
with requests.Session() as s:
counts = collections.Counter(send_completion(s, rh.url) for _ in range(20))
assert len(counts) > 1, f"Random fallback should distribute: {counts}"
def send_completion(session, base_url, routing_key=None):
headers = {ROUTING_KEY_HEADER: routing_key} if routing_key is not None else {}
r = session.post(
f"{base_url}/v1/completions",
json={"model": "test", "prompt": "hi", "max_tokens": 1, "stream": False},
headers=headers,
)
assert r.status_code == 200
return r.headers.get("X-Worker-Id") or r.json().get("worker_id")
-8
View File
@@ -359,10 +359,6 @@ pub struct ChatCompletionRequest {
/// Random seed for sampling for deterministic outputs /// Random seed for sampling for deterministic outputs
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub sampling_seed: Option<u64>, pub sampling_seed: Option<u64>,
/// Routing ID for manual routing policy
#[serde(skip_serializing_if = "Option::is_none")]
pub routing_id: Option<String>,
} }
// ============================================================================ // ============================================================================
@@ -700,10 +696,6 @@ impl GenerationRequest for ChatCompletionRequest {
buffer buffer
} }
fn get_routing_id(&self) -> Option<&str> {
self.routing_id.as_deref()
}
} }
// ============================================================================ // ============================================================================
@@ -30,10 +30,6 @@ pub struct ClassifyRequest {
/// SGLang extension: request id for tracking /// SGLang extension: request id for tracking
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub rid: Option<String>, pub rid: Option<String>,
/// Routing ID for manual routing policy
#[serde(skip_serializing_if = "Option::is_none")]
pub routing_id: Option<String>,
} }
impl GenerationRequest for ClassifyRequest { impl GenerationRequest for ClassifyRequest {
@@ -58,8 +54,4 @@ impl GenerationRequest for ClassifyRequest {
_ => String::new(), _ => String::new(),
} }
} }
fn get_routing_id(&self) -> Option<&str> {
self.routing_id.as_deref()
}
} }
@@ -36,9 +36,6 @@ pub trait GenerationRequest: Send + Sync {
/// Extract text content for routing decisions /// Extract text content for routing decisions
fn extract_text_for_routing(&self) -> String; fn extract_text_for_routing(&self) -> String;
/// Get routing ID for manual routing policy
fn get_routing_id(&self) -> Option<&str>;
} }
// ============================================================================ // ============================================================================
@@ -145,10 +145,6 @@ pub struct CompletionRequest {
/// Additional fields including bootstrap info for PD routing /// Additional fields including bootstrap info for PD routing
#[serde(flatten)] #[serde(flatten)]
pub other: Map<String, Value>, pub other: Map<String, Value>,
/// Routing ID for manual routing policy
#[serde(skip_serializing_if = "Option::is_none")]
pub routing_id: Option<String>,
} }
impl GenerationRequest for CompletionRequest { impl GenerationRequest for CompletionRequest {
@@ -166,10 +162,6 @@ impl GenerationRequest for CompletionRequest {
StringOrArray::Array(v) => v.join(" "), StringOrArray::Array(v) => v.join(" "),
} }
} }
fn get_routing_id(&self) -> Option<&str> {
self.routing_id.as_deref()
}
} }
// ============================================================================ // ============================================================================
@@ -31,10 +31,6 @@ pub struct EmbeddingRequest {
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub rid: Option<String>, pub rid: Option<String>,
/// Routing ID for manual routing policy
#[serde(skip_serializing_if = "Option::is_none")]
pub routing_id: Option<String>,
/// SGLang extension: enable/disable logging of metrics for this request /// SGLang extension: enable/disable logging of metrics for this request
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub log_metrics: Option<bool>, pub log_metrics: Option<bool>,
@@ -62,10 +58,6 @@ impl GenerationRequest for EmbeddingRequest {
_ => String::new(), _ => String::new(),
} }
} }
fn get_routing_id(&self) -> Option<&str> {
self.routing_id.as_deref()
}
} }
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
@@ -167,10 +167,6 @@ pub struct GenerateRequest {
/// Request ID for tracking (inherited from BaseReq in Python) /// Request ID for tracking (inherited from BaseReq in Python)
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub rid: Option<String>, pub rid: Option<String>,
/// Routing ID for manual routing policy
#[serde(skip_serializing_if = "Option::is_none")]
pub routing_id: Option<String>,
} }
impl Normalizable for GenerateRequest { impl Normalizable for GenerateRequest {
@@ -239,10 +235,6 @@ impl GenerationRequest for GenerateRequest {
// No text input found // No text input found
String::new() String::new()
} }
fn get_routing_id(&self) -> Option<&str> {
self.routing_id.as_deref()
}
} }
// ============================================================================ // ============================================================================
@@ -52,10 +52,6 @@ pub struct RerankRequest {
/// User identifier /// User identifier
pub user: Option<String>, pub user: Option<String>,
/// Routing ID for manual routing policy
#[serde(skip_serializing_if = "Option::is_none")]
pub routing_id: Option<String>,
} }
impl GenerationRequest for RerankRequest { impl GenerationRequest for RerankRequest {
@@ -70,10 +66,6 @@ impl GenerationRequest for RerankRequest {
fn extract_text_for_routing(&self) -> String { fn extract_text_for_routing(&self) -> String {
self.query.clone() self.query.clone()
} }
fn get_routing_id(&self) -> Option<&str> {
self.routing_id.as_deref()
}
} }
impl super::validated::Normalizable for RerankRequest { impl super::validated::Normalizable for RerankRequest {
@@ -215,7 +207,6 @@ impl From<V1RerankReqInput> for RerankRequest {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
} }
} }
} }
@@ -616,10 +616,6 @@ pub struct ResponsesRequest {
#[serde(default = "default_repetition_penalty")] #[serde(default = "default_repetition_penalty")]
#[validate(range(min = 0.0, max = 2.0))] #[validate(range(min = 0.0, max = 2.0))]
pub repetition_penalty: f32, pub repetition_penalty: f32,
/// Routing ID for manual routing policy
#[serde(skip_serializing_if = "Option::is_none")]
pub routing_id: Option<String>,
} }
#[derive(Debug, Clone, Deserialize, Serialize)] #[derive(Debug, Clone, Deserialize, Serialize)]
@@ -663,7 +659,6 @@ impl Default for ResponsesRequest {
top_k: default_top_k(), top_k: default_top_k(),
min_p: 0.0, min_p: 0.0,
repetition_penalty: default_repetition_penalty(), repetition_penalty: default_repetition_penalty(),
routing_id: None,
} }
} }
} }
@@ -775,10 +770,6 @@ impl GenerationRequest for ResponsesRequest {
.join(" "), .join(" "),
} }
} }
fn get_routing_id(&self) -> Option<&str> {
self.routing_id.as_deref()
}
} }
/// Validate conversation ID format /// Validate conversation ID format
@@ -19,6 +19,7 @@ use crate::{
context::{PreparationOutput, RequestContext, RequestType}, context::{PreparationOutput, RequestContext, RequestType},
utils, utils,
}, },
header_utils,
}, },
}; };
@@ -123,7 +124,7 @@ impl HarmonyPreparationStage {
// Step 4: Store results // Step 4: Store results
ctx.state.preparation = Some(PreparationOutput { ctx.state.preparation = Some(PreparationOutput {
original_text: None, original_text: None,
routing_id: request.routing_id.clone(), routing_id: header_utils::extract_routing_id(ctx.input.headers.as_ref()),
token_ids: build_output.input_ids, token_ids: build_output.input_ids,
processed_messages: None, processed_messages: None,
tool_constraints, tool_constraints,
@@ -204,7 +205,7 @@ impl HarmonyPreparationStage {
// Step 4: Store results with constraint // Step 4: Store results with constraint
ctx.state.preparation = Some(PreparationOutput { ctx.state.preparation = Some(PreparationOutput {
original_text: None, original_text: None,
routing_id: request.routing_id.clone(), routing_id: header_utils::extract_routing_id(ctx.input.headers.as_ref()),
token_ids: build_output.input_ids, token_ids: build_output.input_ids,
processed_messages: None, processed_messages: None,
tool_constraints: constraint, tool_constraints: constraint,
@@ -492,7 +492,6 @@ pub(super) async fn execute_tool_loop(
top_k: current_request.top_k, top_k: current_request.top_k,
min_p: current_request.min_p, min_p: current_request.min_p,
repetition_penalty: current_request.repetition_penalty, repetition_penalty: current_request.repetition_penalty,
routing_id: current_request.routing_id.clone(),
}; };
// Continue to next iteration // Continue to next iteration
@@ -1071,7 +1070,6 @@ async fn execute_tool_loop_streaming_internal(
top_k: current_request.top_k, top_k: current_request.top_k,
min_p: current_request.min_p, min_p: current_request.min_p,
repetition_penalty: current_request.repetition_penalty, repetition_penalty: current_request.repetition_penalty,
routing_id: current_request.routing_id.clone(),
}; };
continue; continue;
@@ -15,6 +15,7 @@ use crate::{
context::{PreparationOutput, RequestContext}, context::{PreparationOutput, RequestContext},
utils, utils,
}, },
header_utils,
}, },
}; };
@@ -96,7 +97,7 @@ impl ChatPreparationStage {
// Store results in context // Store results in context
ctx.state.preparation = Some(PreparationOutput { ctx.state.preparation = Some(PreparationOutput {
original_text: Some(processed_messages.text.clone()), original_text: Some(processed_messages.text.clone()),
routing_id: request.routing_id.clone(), routing_id: header_utils::extract_routing_id(ctx.input.headers.as_ref()),
token_ids, token_ids,
processed_messages: Some(processed_messages), processed_messages: Some(processed_messages),
tool_constraints: tool_call_constraint, tool_constraints: tool_call_constraint,
@@ -13,6 +13,7 @@ use crate::{
context::{PreparationOutput, RequestContext, RequestType}, context::{PreparationOutput, RequestContext, RequestType},
utils, utils,
}, },
header_utils,
}, },
}; };
@@ -47,9 +48,9 @@ impl PipelineStage for EmbeddingPreparationStage {
)); ));
}; };
// Extract text and routing_id from request before borrowing ctx mutably // Extract text from request before borrowing ctx mutably
let text = request.extract_text_for_routing(); let text = request.extract_text_for_routing();
let routing_id = request.routing_id.clone(); let routing_id = header_utils::extract_routing_id(ctx.input.headers.as_ref());
if text.is_empty() { if text.is_empty() {
return Err(error::bad_request( return Err(error::bad_request(
"empty_input", "empty_input",
@@ -15,6 +15,7 @@ use crate::{
context::{PreparationOutput, RequestContext}, context::{PreparationOutput, RequestContext},
utils, utils,
}, },
header_utils,
}, },
tokenizer::traits::Tokenizer, tokenizer::traits::Tokenizer,
}; };
@@ -68,7 +69,7 @@ impl GeneratePreparationStage {
ctx.state.preparation = Some(PreparationOutput { ctx.state.preparation = Some(PreparationOutput {
original_text, original_text,
routing_id: request.routing_id.clone(), routing_id: header_utils::extract_routing_id(ctx.input.headers.as_ref()),
token_ids, token_ids,
processed_messages: None, processed_messages: None,
tool_constraints: None, tool_constraints: None,
@@ -157,6 +157,17 @@ pub fn apply_provider_headers(
req req
} }
/// Header name for routing key used by manual routing policy
pub const ROUTING_KEY_HEADER: &str = "X-SMG-Routing-Key";
/// Extract routing ID from HTTP headers for manual routing policy
pub fn extract_routing_id(headers: Option<&HeaderMap>) -> Option<String> {
headers
.and_then(|h| h.get(ROUTING_KEY_HEADER))
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
}
/// Extract auth header with passthrough semantics. /// Extract auth header with passthrough semantics.
/// ///
/// Passthrough mode: User's Authorization header takes priority. /// Passthrough mode: User's Authorization header takes priority.
@@ -183,3 +194,31 @@ pub fn extract_auth_header(
.and_then(|k| HeaderValue::from_str(&format!("Bearer {}", k)).ok()) .and_then(|k| HeaderValue::from_str(&format!("Bearer {}", k)).ok())
}) })
} }
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_extract_routing_id_with_header() {
let mut headers = HeaderMap::new();
headers.insert(ROUTING_KEY_HEADER, HeaderValue::from_static("user-123"));
let result = extract_routing_id(Some(&headers));
assert_eq!(result, Some("user-123".to_string()));
}
#[test]
fn test_extract_routing_id_without_header() {
let headers = HeaderMap::new();
let result = extract_routing_id(Some(&headers));
assert_eq!(result, None);
}
#[test]
fn test_extract_routing_id_empty_value() {
let mut headers = HeaderMap::new();
headers.insert(ROUTING_KEY_HEADER, HeaderValue::from_static(""));
let result = extract_routing_id(Some(&headers));
assert_eq!(result, Some("".to_string()));
}
}
@@ -1245,7 +1245,7 @@ impl RouterTrait for PDRouter {
is_stream, is_stream,
return_logprob, return_logprob,
request_text, request_text,
routing_id: body.routing_id.clone(), routing_id: header_utils::extract_routing_id(headers),
model_id, model_id,
}; };
@@ -1287,7 +1287,7 @@ impl RouterTrait for PDRouter {
is_stream, is_stream,
return_logprob, return_logprob,
request_text, request_text,
routing_id: body.routing_id.clone(), routing_id: header_utils::extract_routing_id(headers),
model_id, model_id,
}; };
@@ -1321,7 +1321,7 @@ impl RouterTrait for PDRouter {
is_stream, is_stream,
return_logprob, return_logprob,
request_text, request_text,
routing_id: body.routing_id.clone(), routing_id: header_utils::extract_routing_id(headers),
model_id, model_id,
}; };
@@ -1346,7 +1346,7 @@ impl RouterTrait for PDRouter {
is_stream: false, is_stream: false,
return_logprob: false, return_logprob: false,
request_text: req_text, request_text: req_text,
routing_id: body.routing_id.clone(), routing_id: header_utils::extract_routing_id(headers),
model_id, model_id,
}; };
+1 -1
View File
@@ -191,7 +191,7 @@ impl Router {
let start = Instant::now(); let start = Instant::now();
let is_stream = typed_req.is_stream(); let is_stream = typed_req.is_stream();
let text = typed_req.extract_text_for_routing(); let text = typed_req.extract_text_for_routing();
let routing_id = typed_req.get_routing_id().map(|s| s.to_string()); let routing_id = header_utils::extract_routing_id(headers);
let info = crate::policies::SelectWorkerInfo { let info = crate::policies::SelectWorkerInfo {
request_text: Some(&text), request_text: Some(&text),
routing_id: routing_id.as_deref(), routing_id: routing_id.as_deref(),
@@ -105,7 +105,6 @@ async fn test_non_streaming_mcp_minimal_e2e_with_persistence() {
min_p: 0.0, min_p: 0.0,
repetition_penalty: 1.0, repetition_penalty: 1.0,
conversation: None, conversation: None,
routing_id: None,
}; };
let resp = router let resp = router
@@ -329,7 +328,6 @@ fn test_responses_request_creation() {
min_p: 0.0, min_p: 0.0,
repetition_penalty: 1.0, repetition_penalty: 1.0,
conversation: None, conversation: None,
routing_id: None,
}; };
assert!(!request.is_stream()); assert!(!request.is_stream());
@@ -374,7 +372,6 @@ fn test_responses_request_sglang_extensions() {
min_p: 0.05, min_p: 0.05,
repetition_penalty: 1.1, repetition_penalty: 1.1,
conversation: None, conversation: None,
routing_id: None,
}; };
// Verify SGLang extensions are present // Verify SGLang extensions are present
@@ -490,7 +487,6 @@ fn test_json_serialization() {
min_p: 0.1, min_p: 0.1,
repetition_penalty: 1.2, repetition_penalty: 1.2,
conversation: None, conversation: None,
routing_id: None,
}; };
let json = serde_json::to_string(&request).expect("Serialization should work"); let json = serde_json::to_string(&request).expect("Serialization should work");
@@ -597,7 +593,6 @@ async fn test_multi_turn_loop_with_mcp() {
min_p: 0.0, min_p: 0.0,
repetition_penalty: 1.0, repetition_penalty: 1.0,
conversation: None, conversation: None,
routing_id: None,
}; };
// Execute the request (this should trigger the multi-turn loop) // Execute the request (this should trigger the multi-turn loop)
@@ -747,7 +742,6 @@ async fn test_max_tool_calls_limit() {
min_p: 0.0, min_p: 0.0,
repetition_penalty: 1.0, repetition_penalty: 1.0,
conversation: None, conversation: None,
routing_id: None,
}; };
let response = router.route_responses(None, &req, None).await; let response = router.route_responses(None, &req, None).await;
@@ -920,7 +914,6 @@ async fn test_streaming_with_mcp_tool_calls() {
min_p: 0.0, min_p: 0.0,
repetition_penalty: 1.0, repetition_penalty: 1.0,
conversation: None, conversation: None,
routing_id: None,
}; };
let response = router.route_responses(None, &req, None).await; let response = router.route_responses(None, &req, None).await;
@@ -1201,7 +1194,6 @@ async fn test_streaming_multi_turn_with_mcp() {
min_p: 0.0, min_p: 0.0,
repetition_penalty: 1.0, repetition_penalty: 1.0,
conversation: None, conversation: None,
routing_id: None,
}; };
let response = router.route_responses(None, &req, None).await; let response = router.route_responses(None, &req, None).await;
@@ -10,7 +10,6 @@ fn test_embedding_request_serialization_string_input() {
user: Some("user-1".to_string()), user: Some("user-1".to_string()),
dimensions: Some(128), dimensions: Some(128),
rid: Some("rid-123".to_string()), rid: Some("rid-123".to_string()),
routing_id: None,
log_metrics: None, log_metrics: None,
}; };
@@ -34,7 +33,6 @@ fn test_embedding_request_serialization_array_input() {
user: None, user: None,
dimensions: None, dimensions: None,
rid: None, rid: None,
routing_id: None,
log_metrics: None, log_metrics: None,
}; };
@@ -53,7 +51,6 @@ fn test_embedding_generation_request_trait_string() {
user: None, user: None,
dimensions: None, dimensions: None,
rid: None, rid: None,
routing_id: None,
log_metrics: None, log_metrics: None,
}; };
assert!(!req.is_stream()); assert!(!req.is_stream());
@@ -70,7 +67,6 @@ fn test_embedding_generation_request_trait_array() {
user: None, user: None,
dimensions: None, dimensions: None,
rid: None, rid: None,
routing_id: None,
log_metrics: None, log_metrics: None,
}; };
assert_eq!(req.extract_text_for_routing(), "hello world"); assert_eq!(req.extract_text_for_routing(), "hello world");
@@ -85,7 +81,6 @@ fn test_embedding_generation_request_trait_non_text() {
user: None, user: None,
dimensions: None, dimensions: None,
rid: None, rid: None,
routing_id: None,
log_metrics: None, log_metrics: None,
}; };
assert_eq!(req.extract_text_for_routing(), ""); assert_eq!(req.extract_text_for_routing(), "");
@@ -100,7 +95,6 @@ fn test_embedding_generation_request_trait_mixed_array_ignores_nested() {
user: None, user: None,
dimensions: None, dimensions: None,
rid: None, rid: None,
routing_id: None,
log_metrics: None, log_metrics: None,
}; };
// Only top-level string elements are extracted // Only top-level string elements are extracted
-15
View File
@@ -17,7 +17,6 @@ fn test_rerank_request_serialization() {
return_documents: true, return_documents: true,
rid: Some(StringOrArray::String("req-123".to_string())), rid: Some(StringOrArray::String("req-123".to_string())),
user: Some("user-456".to_string()), user: Some("user-456".to_string()),
routing_id: None,
}; };
let serialized = to_string(&request).unwrap(); let serialized = to_string(&request).unwrap();
@@ -60,7 +59,6 @@ fn test_rerank_request_validation_success() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
assert!(request.validate().is_ok()); assert!(request.validate().is_ok());
@@ -76,7 +74,6 @@ fn test_rerank_request_validation_empty_query() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
let result = request.validate(); let result = request.validate();
@@ -93,7 +90,6 @@ fn test_rerank_request_validation_whitespace_query() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
let result = request.validate(); let result = request.validate();
@@ -110,7 +106,6 @@ fn test_rerank_request_validation_empty_documents() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
let result = request.validate(); let result = request.validate();
@@ -127,7 +122,6 @@ fn test_rerank_request_validation_top_k_zero() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
let result = request.validate(); let result = request.validate();
@@ -144,7 +138,6 @@ fn test_rerank_request_validation_top_k_greater_than_docs() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
// This should pass but log a warning // This should pass but log a warning
@@ -161,7 +154,6 @@ fn test_rerank_request_effective_top_k() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
assert_eq!(request.effective_top_k(), 2); assert_eq!(request.effective_top_k(), 2);
@@ -177,7 +169,6 @@ fn test_rerank_request_effective_top_k_none() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
assert_eq!(request.effective_top_k(), 3); assert_eq!(request.effective_top_k(), 3);
@@ -399,7 +390,6 @@ fn test_rerank_request_generation_request_trait() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
assert_eq!(request.get_model(), Some("test-model")); assert_eq!(request.get_model(), Some("test-model"));
@@ -418,7 +408,6 @@ fn test_rerank_request_very_long_query() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
assert!(request.validate().is_ok()); assert!(request.validate().is_ok());
@@ -435,7 +424,6 @@ fn test_rerank_request_many_documents() {
return_documents: true, return_documents: true,
rid: None, rid: None,
user: None, user: None,
routing_id: None,
}; };
assert!(request.validate().is_ok()); assert!(request.validate().is_ok());
@@ -455,7 +443,6 @@ fn test_rerank_request_special_characters() {
return_documents: true, return_documents: true,
rid: Some(StringOrArray::String("req-🚀-123".to_string())), rid: Some(StringOrArray::String("req-🚀-123".to_string())),
user: Some("user-🎉-456".to_string()), user: Some("user-🎉-456".to_string()),
routing_id: None,
}; };
assert!(request.validate().is_ok()); assert!(request.validate().is_ok());
@@ -474,7 +461,6 @@ fn test_rerank_request_rid_array() {
"req2".to_string(), "req2".to_string(),
])), ])),
user: None, user: None,
routing_id: None,
}; };
assert!(request.validate().is_ok()); assert!(request.validate().is_ok());
@@ -529,7 +515,6 @@ fn test_full_rerank_workflow() {
return_documents: true, return_documents: true,
rid: Some(StringOrArray::String("req-123".to_string())), rid: Some(StringOrArray::String("req-123".to_string())),
user: Some("user-456".to_string()), user: Some("user-456".to_string()),
routing_id: None,
}; };
// Validate request // Validate request
@@ -89,7 +89,6 @@ fn create_minimal_completion_request() -> CompletionRequest {
return_hidden_states: false, return_hidden_states: false,
sampling_seed: None, sampling_seed: None,
other: serde_json::Map::new(), other: serde_json::Map::new(),
routing_id: None,
} }
} }
@@ -640,7 +639,6 @@ async fn test_unsupported_endpoints() {
return_bytes: false, return_bytes: false,
return_entropy: false, return_entropy: false,
rid: None, rid: None,
routing_id: None,
}; };
let response = router.route_generate(None, &generate_request, None).await; let response = router.route_generate(None, &generate_request, None).await;