Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
113 changes: 45 additions & 68 deletions model_gateway/src/routers/grpc/common/stages/client_acquisition.rs
Original file line number Diff line number Diff line change
@@ -1,92 +1,69 @@
//! Client acquisition stage: Get gRPC clients from selected workers
//! Client acquisition: get gRPC clients from selected workers.
//!
//! Called once at ingress and again per retry attempt after worker
//! re-selection.

use std::sync::Arc;

use async_trait::async_trait;
use axum::response::Response;
use tracing::error;

use super::PipelineStage;
use crate::{
routers::{
common::overload,
error,
grpc::{
backend_client::BackendClient,
context::{ClientSelection, RequestContext, WorkerSelection},
context::{ClientSelection, WorkerSelection},
},
},
worker::Worker,
};

/// Client acquisition stage: Get gRPC clients from selected workers
pub(crate) struct ClientAcquisitionStage;

#[async_trait]
impl PipelineStage for ClientAcquisitionStage {
async fn execute(&self, ctx: &mut RequestContext) -> Result<Option<Response>, Response> {
let workers = ctx.state.workers.as_ref().ok_or_else(|| {
error!(
function = "ClientAcquisitionStage::execute",
"Worker selection stage not completed"
);
error::internal_error(
"worker_selection_not_completed",
"Worker selection not completed",
)
})?;

// Dispatch-time re-check: one relaxed atomic read per already-chosen
// worker, closing the window between selection and dispatch in which a
// load report can flip the veto.
let model_id = ctx.input.model_id.as_str();
let clients = match workers {
WorkerSelection::Single { worker } => {
if let Some(shed) = overload::shed_if_worker_overloaded(worker.as_ref(), model_id) {
return Err(shed);
}
let client = get_backend_client_from_worker(worker).await?;
ClientSelection::Single { client }
/// Acquire backend clients for the selected workers, with a dispatch-time
/// overload re-check: one relaxed atomic read per already-chosen worker,
/// closing the window between selection and dispatch in which a load report
/// can flip the veto.
pub(crate) async fn acquire_clients(
workers: &WorkerSelection,
model_id: &str,
) -> Result<ClientSelection, Response> {
match workers {
WorkerSelection::Single { worker } => {
if let Some(shed) = overload::shed_if_worker_overloaded(worker.as_ref(), model_id) {
return Err(shed);
}
WorkerSelection::Disaggregated {
encode_assignments,
prefill,
decode,
..
} => {
// Every assigned leg, encode included: an encode worker is
// vetoed at selection through the same filter, so leaving it out
// of the re-check would be the one dispatch path that can send
// to a worker known to be over the ceiling.
if let Some(shed) = overload::shed_if_worker_overloaded(prefill.as_ref(), model_id)
.or_else(|| overload::shed_if_worker_overloaded(decode.as_ref(), model_id))
.or_else(|| {
encode_assignments.iter().flatten().find_map(|assignment| {
overload::shed_if_worker_overloaded(
assignment.worker.as_ref(),
model_id,
)
})
let client = get_backend_client_from_worker(worker).await?;
Ok(ClientSelection::Single { client })
}
WorkerSelection::Disaggregated {
encode_assignments,
prefill,
decode,
..
} => {
// Every assigned leg, encode included: an encode worker is
// vetoed at selection through the same filter, so leaving it out
// of the re-check would be the one dispatch path that can send
// to a worker known to be over the ceiling.
if let Some(shed) = overload::shed_if_worker_overloaded(prefill.as_ref(), model_id)
.or_else(|| overload::shed_if_worker_overloaded(decode.as_ref(), model_id))
.or_else(|| {
encode_assignments.iter().flatten().find_map(|assignment| {
overload::shed_if_worker_overloaded(assignment.worker.as_ref(), model_id)
})
{
return Err(shed);
}
let prefill_client = get_backend_client_from_worker(prefill).await?;
let decode_client = get_backend_client_from_worker(decode).await?;

ClientSelection::Disaggregated {
prefill: prefill_client,
decode: decode_client,
}
})
{
return Err(shed);
}
};

ctx.state.clients = Some(clients);
Ok(None)
}
let prefill_client = get_backend_client_from_worker(prefill).await?;
let decode_client = get_backend_client_from_worker(decode).await?;

fn name(&self) -> &'static str {
"ClientAcquisition"
Ok(ClientSelection::Disaggregated {
prefill: prefill_client,
decode: decode_client,
})
}
}
}

Expand Down
100 changes: 30 additions & 70 deletions model_gateway/src/routers/grpc/common/stages/dispatch_metadata.rs
Original file line number Diff line number Diff line change
@@ -1,75 +1,35 @@
//! Dispatch metadata stage: Prepare metadata for dispatch
//! Dispatch metadata: per-attempt response metadata derived from the stamped
//! plan and the attempt's worker selection.

use std::time::{SystemTime, UNIX_EPOCH};

use async_trait::async_trait;
use axum::response::Response;
use tracing::error;

use super::PipelineStage;
use crate::routers::{
error,
grpc::context::{DispatchMetadata, RequestContext, RequestType, WorkerSelection},
};

/// Dispatch metadata stage: Prepare metadata for dispatch
pub(crate) struct DispatchMetadataStage;

#[async_trait]
impl PipelineStage for DispatchMetadataStage {
async fn execute(&self, ctx: &mut RequestContext) -> Result<Option<Response>, Response> {
let execution_plan = ctx.state.execution_plan.as_ref().ok_or_else(|| {
error!(
function = "DispatchMetadataStage::execute",
"Execution plan not built"
);
error::internal_error("execution_plan_not_built", "Execution plan not built")
})?;

let request_id = execution_plan.request_id().to_string();
// The model the response reports. `RequestContext::new` already
// rewrote every one of these to the canonical model ID, so a request
// that arrived under an alias is answered under the canonical name.
let model = match &ctx.input.request_type {
RequestType::Chat(req) => req.model.clone(),
RequestType::Completion(req) => req.model.clone(),
// `GenerateRequest` carries a model field too, but callers of the
// native `/generate` route may leave it empty, so prefer the
// model the router resolved.
RequestType::Generate(_req) => ctx.input.model_id.clone(),
RequestType::Responses(req) => req.model.clone(),
RequestType::Embedding(req) => req.model.clone(),
RequestType::Classify(req) => req.model.clone(),
RequestType::Messages(req) => req.model.clone(),
};

let weight_version = ctx
.state
.workers
.as_ref()
.map(|w| match w {
WorkerSelection::Single { worker } => worker,
WorkerSelection::Disaggregated { decode, .. } => decode,
})
.and_then(|w| w.metadata().spec.labels.get("weight_version").cloned())
.unwrap_or_else(|| "default".to_string());

let created = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();

ctx.state.dispatch = Some(DispatchMetadata {
request_id,
model,
created,
weight_version: Some(weight_version),
});

Ok(None)
}

fn name(&self) -> &'static str {
"DispatchMetadata"
use crate::routers::grpc::context::{DispatchMetadata, ExecutionPlan, WorkerSelection};

/// Metadata for one dispatch attempt. `dispatch_model` was captured from the
/// request at the build boundary (already canonical); the weight version
/// comes from the attempt's selected worker.
pub(crate) fn prepare_dispatch_metadata(
plan: &ExecutionPlan,
dispatch_model: &str,
workers: Option<&WorkerSelection>,
) -> DispatchMetadata {
let weight_version = workers
.map(|w| match w {
WorkerSelection::Single { worker } => worker,
WorkerSelection::Disaggregated { decode, .. } => decode,
})
.and_then(|w| w.metadata().spec.labels.get("weight_version").cloned())
.unwrap_or_else(|| "default".to_string());

let created = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();

DispatchMetadata {
request_id: plan.request_id().to_string(),
model: dispatch_model.to_string(),
created,
weight_version: Some(weight_version),
}
}
15 changes: 5 additions & 10 deletions model_gateway/src/routers/grpc/common/stages/encode.rs
Original file line number Diff line number Diff line change
Expand Up @@ -58,19 +58,19 @@ impl EncodeStage {

#[async_trait]
impl PipelineStage for EncodeStage {
async fn execute(&self, ctx: &mut RequestContext) -> Result<Option<Response>, Response> {
async fn execute(&self, ctx: &mut RequestContext) -> Result<(), Response> {
if ctx
.state
.workers
.as_ref()
.and_then(WorkerSelection::encode_assignments)
.is_none_or(|assignments| assignments.is_empty())
{
return Ok(None);
return Ok(());
}

let Some(intermediate) = ctx.state.multimodal_intermediate.as_ref() else {
return Ok(None);
return Ok(());
};

let plan = build_plan(
Expand All @@ -87,25 +87,20 @@ impl PipelineStage for EncodeStage {
// request building takes the plain prefill path (with pixels) rather
// than the pixel-drop encode path.
if plan.is_empty() {
return Ok(None);
return Ok(());
}

let (bootstrap_info, dispatch) = plan.into_parts();
ctx.state.encode_outputs = Some(EncodeOutputs {
bootstrap_info,
dispatch,
});
Ok(None)
Ok(())
}

fn name(&self) -> &'static str {
"Encode"
}

#[cfg(test)]
fn signature(&self) -> String {
"EncodeStage".to_string()
}
}

pub(crate) struct EncodePlan {
Expand Down
Loading
Loading