Add manual routing policy for router (#15586)
This commit is contained in:
@@ -359,6 +359,10 @@ pub struct ChatCompletionRequest {
|
||||
/// Random seed for sampling for deterministic outputs
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub sampling_seed: Option<u64>,
|
||||
|
||||
/// Routing ID for manual routing policy
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub routing_id: Option<String>,
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
@@ -696,6 +700,10 @@ impl GenerationRequest for ChatCompletionRequest {
|
||||
|
||||
buffer
|
||||
}
|
||||
|
||||
fn get_routing_id(&self) -> Option<&str> {
|
||||
self.routing_id.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
|
||||
@@ -30,6 +30,10 @@ pub struct ClassifyRequest {
|
||||
/// SGLang extension: request id for tracking
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
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 {
|
||||
@@ -54,4 +58,8 @@ impl GenerationRequest for ClassifyRequest {
|
||||
_ => String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn get_routing_id(&self) -> Option<&str> {
|
||||
self.routing_id.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -36,6 +36,9 @@ pub trait GenerationRequest: Send + Sync {
|
||||
|
||||
/// Extract text content for routing decisions
|
||||
fn extract_text_for_routing(&self) -> String;
|
||||
|
||||
/// Get routing ID for manual routing policy
|
||||
fn get_routing_id(&self) -> Option<&str>;
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
|
||||
@@ -145,6 +145,10 @@ pub struct CompletionRequest {
|
||||
/// Additional fields including bootstrap info for PD routing
|
||||
#[serde(flatten)]
|
||||
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 {
|
||||
@@ -162,6 +166,10 @@ impl GenerationRequest for CompletionRequest {
|
||||
StringOrArray::Array(v) => v.join(" "),
|
||||
}
|
||||
}
|
||||
|
||||
fn get_routing_id(&self) -> Option<&str> {
|
||||
self.routing_id.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
|
||||
@@ -31,6 +31,10 @@ pub struct EmbeddingRequest {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
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
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub log_metrics: Option<bool>,
|
||||
@@ -58,6 +62,10 @@ impl GenerationRequest for EmbeddingRequest {
|
||||
_ => String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn get_routing_id(&self) -> Option<&str> {
|
||||
self.routing_id.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
|
||||
@@ -167,6 +167,10 @@ pub struct GenerateRequest {
|
||||
/// Request ID for tracking (inherited from BaseReq in Python)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
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 {
|
||||
@@ -235,6 +239,10 @@ impl GenerationRequest for GenerateRequest {
|
||||
// No text input found
|
||||
String::new()
|
||||
}
|
||||
|
||||
fn get_routing_id(&self) -> Option<&str> {
|
||||
self.routing_id.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
|
||||
@@ -52,6 +52,10 @@ pub struct RerankRequest {
|
||||
|
||||
/// User identifier
|
||||
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 {
|
||||
@@ -66,6 +70,10 @@ impl GenerationRequest for RerankRequest {
|
||||
fn extract_text_for_routing(&self) -> String {
|
||||
self.query.clone()
|
||||
}
|
||||
|
||||
fn get_routing_id(&self) -> Option<&str> {
|
||||
self.routing_id.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
impl super::validated::Normalizable for RerankRequest {
|
||||
@@ -207,6 +215,7 @@ impl From<V1RerankReqInput> for RerankRequest {
|
||||
return_documents: true,
|
||||
rid: None,
|
||||
user: None,
|
||||
routing_id: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -616,6 +616,10 @@ pub struct ResponsesRequest {
|
||||
#[serde(default = "default_repetition_penalty")]
|
||||
#[validate(range(min = 0.0, max = 2.0))]
|
||||
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)]
|
||||
@@ -659,6 +663,7 @@ impl Default for ResponsesRequest {
|
||||
top_k: default_top_k(),
|
||||
min_p: 0.0,
|
||||
repetition_penalty: default_repetition_penalty(),
|
||||
routing_id: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -770,6 +775,10 @@ impl GenerationRequest for ResponsesRequest {
|
||||
.join(" "),
|
||||
}
|
||||
}
|
||||
|
||||
fn get_routing_id(&self) -> Option<&str> {
|
||||
self.routing_id.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
/// Validate conversation ID format
|
||||
|
||||
Reference in New Issue
Block a user