Add manual routing policy for router (#15586)
This commit is contained in:
@@ -44,6 +44,7 @@ fn test_backward_compatibility_with_empty_model_id() {
|
||||
&workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test request"),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
assert!(selected.is_some(), "Should select a worker");
|
||||
@@ -102,15 +103,24 @@ fn test_mixed_model_ids() {
|
||||
|
||||
let default_workers: Vec<Arc<dyn Worker>> =
|
||||
vec![Arc::new(worker1.clone()), Arc::new(worker3.clone())];
|
||||
let info = SelectWorkerInfo {
|
||||
request_text: Some("test request"),
|
||||
};
|
||||
let selected = policy.select_worker(&default_workers, &info);
|
||||
let selected = policy.select_worker(
|
||||
&default_workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test request"),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
assert!(selected.is_some(), "Should select from default workers");
|
||||
|
||||
let llama_workers: Vec<Arc<dyn Worker>> =
|
||||
vec![Arc::new(worker2.clone()), Arc::new(worker4.clone())];
|
||||
let selected = policy.select_worker(&llama_workers, &info);
|
||||
let selected = policy.select_worker(
|
||||
&llama_workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test request"),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
assert!(selected.is_some(), "Should select from llama-3 workers");
|
||||
|
||||
let all_workers: Vec<Arc<dyn Worker>> = vec![
|
||||
@@ -119,7 +129,13 @@ fn test_mixed_model_ids() {
|
||||
Arc::new(worker3.clone()),
|
||||
Arc::new(worker4.clone()),
|
||||
];
|
||||
let selected = policy.select_worker(&all_workers, &info);
|
||||
let selected = policy.select_worker(
|
||||
&all_workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test request"),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
assert!(selected.is_some(), "Should select from all workers");
|
||||
}
|
||||
|
||||
@@ -156,6 +172,7 @@ fn test_remove_worker_by_url_backward_compat() {
|
||||
&workers,
|
||||
&SelectWorkerInfo {
|
||||
request_text: Some("test"),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
assert_eq!(selected, Some(0), "Should only have worker2 left");
|
||||
|
||||
@@ -105,6 +105,7 @@ async fn test_non_streaming_mcp_minimal_e2e_with_persistence() {
|
||||
min_p: 0.0,
|
||||
repetition_penalty: 1.0,
|
||||
conversation: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let resp = router
|
||||
@@ -328,6 +329,7 @@ fn test_responses_request_creation() {
|
||||
min_p: 0.0,
|
||||
repetition_penalty: 1.0,
|
||||
conversation: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
assert!(!request.is_stream());
|
||||
@@ -372,6 +374,7 @@ fn test_responses_request_sglang_extensions() {
|
||||
min_p: 0.05,
|
||||
repetition_penalty: 1.1,
|
||||
conversation: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
// Verify SGLang extensions are present
|
||||
@@ -487,6 +490,7 @@ fn test_json_serialization() {
|
||||
min_p: 0.1,
|
||||
repetition_penalty: 1.2,
|
||||
conversation: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let json = serde_json::to_string(&request).expect("Serialization should work");
|
||||
@@ -593,6 +597,7 @@ async fn test_multi_turn_loop_with_mcp() {
|
||||
min_p: 0.0,
|
||||
repetition_penalty: 1.0,
|
||||
conversation: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
// Execute the request (this should trigger the multi-turn loop)
|
||||
@@ -742,6 +747,7 @@ async fn test_max_tool_calls_limit() {
|
||||
min_p: 0.0,
|
||||
repetition_penalty: 1.0,
|
||||
conversation: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let response = router.route_responses(None, &req, None).await;
|
||||
@@ -914,6 +920,7 @@ async fn test_streaming_with_mcp_tool_calls() {
|
||||
min_p: 0.0,
|
||||
repetition_penalty: 1.0,
|
||||
conversation: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let response = router.route_responses(None, &req, None).await;
|
||||
@@ -1194,6 +1201,7 @@ async fn test_streaming_multi_turn_with_mcp() {
|
||||
min_p: 0.0,
|
||||
repetition_penalty: 1.0,
|
||||
conversation: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let response = router.route_responses(None, &req, None).await;
|
||||
|
||||
@@ -10,6 +10,7 @@ fn test_embedding_request_serialization_string_input() {
|
||||
user: Some("user-1".to_string()),
|
||||
dimensions: Some(128),
|
||||
rid: Some("rid-123".to_string()),
|
||||
routing_id: None,
|
||||
log_metrics: None,
|
||||
};
|
||||
|
||||
@@ -33,6 +34,7 @@ fn test_embedding_request_serialization_array_input() {
|
||||
user: None,
|
||||
dimensions: None,
|
||||
rid: None,
|
||||
routing_id: None,
|
||||
log_metrics: None,
|
||||
};
|
||||
|
||||
@@ -51,6 +53,7 @@ fn test_embedding_generation_request_trait_string() {
|
||||
user: None,
|
||||
dimensions: None,
|
||||
rid: None,
|
||||
routing_id: None,
|
||||
log_metrics: None,
|
||||
};
|
||||
assert!(!req.is_stream());
|
||||
@@ -67,6 +70,7 @@ fn test_embedding_generation_request_trait_array() {
|
||||
user: None,
|
||||
dimensions: None,
|
||||
rid: None,
|
||||
routing_id: None,
|
||||
log_metrics: None,
|
||||
};
|
||||
assert_eq!(req.extract_text_for_routing(), "hello world");
|
||||
@@ -81,6 +85,7 @@ fn test_embedding_generation_request_trait_non_text() {
|
||||
user: None,
|
||||
dimensions: None,
|
||||
rid: None,
|
||||
routing_id: None,
|
||||
log_metrics: None,
|
||||
};
|
||||
assert_eq!(req.extract_text_for_routing(), "");
|
||||
@@ -95,6 +100,7 @@ fn test_embedding_generation_request_trait_mixed_array_ignores_nested() {
|
||||
user: None,
|
||||
dimensions: None,
|
||||
rid: None,
|
||||
routing_id: None,
|
||||
log_metrics: None,
|
||||
};
|
||||
// Only top-level string elements are extracted
|
||||
|
||||
@@ -17,6 +17,7 @@ fn test_rerank_request_serialization() {
|
||||
return_documents: true,
|
||||
rid: Some(StringOrArray::String("req-123".to_string())),
|
||||
user: Some("user-456".to_string()),
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let serialized = to_string(&request).unwrap();
|
||||
@@ -59,6 +60,7 @@ fn test_rerank_request_validation_success() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
assert!(request.validate().is_ok());
|
||||
@@ -74,6 +76,7 @@ fn test_rerank_request_validation_empty_query() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let result = request.validate();
|
||||
@@ -90,6 +93,7 @@ fn test_rerank_request_validation_whitespace_query() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let result = request.validate();
|
||||
@@ -106,6 +110,7 @@ fn test_rerank_request_validation_empty_documents() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let result = request.validate();
|
||||
@@ -122,6 +127,7 @@ fn test_rerank_request_validation_top_k_zero() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let result = request.validate();
|
||||
@@ -138,6 +144,7 @@ fn test_rerank_request_validation_top_k_greater_than_docs() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
// This should pass but log a warning
|
||||
@@ -154,6 +161,7 @@ fn test_rerank_request_effective_top_k() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
assert_eq!(request.effective_top_k(), 2);
|
||||
@@ -169,6 +177,7 @@ fn test_rerank_request_effective_top_k_none() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
assert_eq!(request.effective_top_k(), 3);
|
||||
@@ -390,6 +399,7 @@ fn test_rerank_request_generation_request_trait() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
assert_eq!(request.get_model(), Some("test-model"));
|
||||
@@ -408,6 +418,7 @@ fn test_rerank_request_very_long_query() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
assert!(request.validate().is_ok());
|
||||
@@ -424,6 +435,7 @@ fn test_rerank_request_many_documents() {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
assert!(request.validate().is_ok());
|
||||
@@ -443,6 +455,7 @@ fn test_rerank_request_special_characters() {
|
||||
return_documents: true,
|
||||
rid: Some(StringOrArray::String("req-🚀-123".to_string())),
|
||||
user: Some("user-🎉-456".to_string()),
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
assert!(request.validate().is_ok());
|
||||
@@ -461,6 +474,7 @@ fn test_rerank_request_rid_array() {
|
||||
"req2".to_string(),
|
||||
])),
|
||||
user: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
assert!(request.validate().is_ok());
|
||||
@@ -515,6 +529,7 @@ fn test_full_rerank_workflow() {
|
||||
return_documents: true,
|
||||
rid: Some(StringOrArray::String("req-123".to_string())),
|
||||
user: Some("user-456".to_string()),
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
// Validate request
|
||||
|
||||
@@ -89,6 +89,7 @@ fn create_minimal_completion_request() -> CompletionRequest {
|
||||
return_hidden_states: false,
|
||||
sampling_seed: None,
|
||||
other: serde_json::Map::new(),
|
||||
routing_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -639,6 +640,7 @@ async fn test_unsupported_endpoints() {
|
||||
return_bytes: false,
|
||||
return_entropy: false,
|
||||
rid: None,
|
||||
routing_id: None,
|
||||
};
|
||||
|
||||
let response = router.route_generate(None, &generate_request, None).await;
|
||||
|
||||
Reference in New Issue
Block a user