Skip to content
Merged
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
2 changes: 2 additions & 0 deletions lib/sidecar/vllm/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,8 @@ pool. Override them with `--grpc-connect-attempt-timeout-secs`,
`--grpc-retry-interval-secs`, and `--grpc-startup-deadline-secs`, or with the
corresponding `DYN_SIDECAR_GRPC_*` environment variables.

Each request owns its response stream but borrows a channel from the shared pool. Aggregate and prefill cancellation drops only that request's stream. Decode cancellation first submits the decode request and retains its stream until the first output token or a response containing `finish_info`, so a NIXL receiver can complete and release the transferred KV; it then drops the stream. If the stream ends early, returns a gRPC error, or produces an invalid response after cancellation, the sidecar logs the failure and reports the request as cancelled. vLLM automatically aborts the corresponding engine request while the pooled HTTP/2 connection remains available to other requests. The sidecar does not call the Control `Abort` RPC.

## Test without vLLM or a GPU

Use the CPU-only `dynamo-vllm-mocker-server` to exercise the same Inference, Control, and health contracts:
Expand Down
127 changes: 96 additions & 31 deletions lib/sidecar/vllm/src/engine.rs
Original file line number Diff line number Diff line change
Expand Up @@ -191,60 +191,125 @@ impl LLMEngine for VllmSidecarEngine {
let mut state = ResponseState::new(&request, self.mode);
let mut proto_request = build_generate_request(request, request_id, self.mode)?;
proto_request.model.clone_from(&self.model.served_name);
let defer_request_cancellation = self.mode.is_decode();
let stopped_ctx = ctx.inner_arc();
let shutdown = self.cancel.clone();
let mut cancellation = Box::pin(async move {
let mut request_cancellation = Box::pin(async move { stopped_ctx.stopped().await });
let mut shutdown_cancellation = Box::pin(async move { shutdown.cancelled().await });
let stream = if defer_request_cancellation {
// Decode must reach vLLM so NIXL can release transferred KV.
tokio::select! {
_ = stopped_ctx.stopped() => {}
_ = shutdown.cancelled() => {}
biased;
_ = shutdown_cancellation.as_mut() => None,
result = client.generate_stream(proto_request) => Some(result?),
}
} else {
tokio::select! {
biased;
_ = shutdown_cancellation.as_mut() => None,
_ = request_cancellation.as_mut() => None,
result = client.generate_stream(proto_request) => Some(result?),
}
});
let stream = tokio::select! {
biased;
_ = cancellation.as_mut() => None,
result = client.generate_stream(proto_request) => Some(result?),
};
Comment thread
connorcarpenter15 marked this conversation as resolved.
let Some(mut stream) = stream else {
let output = cancelled(&state);
return Ok(Box::pin(futures::stream::once(async move { Ok(output) })));
};

Ok(Box::pin(async_stream::stream! {
let mut request_cancelled = false;
let mut first_token_observed = false;
loop {
tokio::select! {
biased;
_ = cancellation.as_mut() => {
yield Ok(cancelled(&state));
break;
let message = if request_cancelled {
tokio::select! {
biased;
_ = shutdown_cancellation.as_mut() => None,
message = stream.message() => Some(message),
}
} else {
tokio::select! {
biased;
_ = shutdown_cancellation.as_mut() => None,
_ = request_cancellation.as_mut() => {
if defer_request_cancellation && !first_token_observed {
request_cancelled = true;
continue;
}
None
}
message = stream.message() => Some(message),
}
message = stream.message() => {
match message {
Ok(Some(response)) => match state.convert(response) {
Ok(Some(output)) => {
let terminal = output.finish_reason.is_some();
yield Ok(output);
if terminal {
break;
};
Comment thread
coderabbitai[bot] marked this conversation as resolved.

let Some(message) = message else {
yield Ok(cancelled(&state));
break;
};
match message {
Ok(Some(response)) => {
let response_has_token = response
.outputs
.as_ref()
.is_some_and(|output| output.num_tokens > 0);
let transfer_completed = response.outputs.as_ref().is_some_and(|output| {
output.num_tokens > 0 || output.finish_info.is_some()
});
match state.convert(response) {
Ok(Some(output)) => {
first_token_observed |= response_has_token;
if request_cancelled && transfer_completed {
if first_token_observed {
ctx.notify_first_token();
}
yield Ok(cancelled(&state));
break;
}
Ok(None) => {}
Err(error) => {
yield Err(error);
let terminal = output.finish_reason.is_some();
yield Ok(output);
if terminal {
break;
}
},
Ok(None) => {
yield Err(client::protocol_error(
"GenerateStream ended before a terminal response",
));
}
Ok(None) => {}
Err(error) if request_cancelled => {
tracing::warn!(
%error,
"vLLM response conversion failed after request cancellation"
);
yield Ok(cancelled(&state));
break;
}
Err(status) => {
yield Err(client::status_to_dynamo("GenerateStream", status));
Err(error) => {
yield Err(error);
break;
}
}
}
Ok(None) if request_cancelled => {
tracing::warn!(
"vLLM GenerateStream ended before transfer completion after request cancellation"
);
yield Ok(cancelled(&state));
break;
}
Ok(None) => {
yield Err(client::protocol_error(
"GenerateStream ended before a terminal response",
));
break;
}
Err(status) if request_cancelled => {
tracing::warn!(
%status,
"vLLM GenerateStream failed before transfer completion after request cancellation"
);
yield Ok(cancelled(&state));
break;
}
Err(status) => {
yield Err(client::status_to_dynamo("GenerateStream", status));
break;
}
Comment thread
connorcarpenter15 marked this conversation as resolved.
}
}
}))
Expand Down
135 changes: 135 additions & 0 deletions lib/sidecar/vllm/src/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,10 @@ struct FakeVllm {
hang_before_headers: Arc<AtomicBool>,
headers_pending: Arc<AtomicBool>,
release_headers: Arc<Notify>,
hold_before_first_token: Arc<AtomicBool>,
close_before_first_token: Arc<AtomicBool>,
first_token_pending: Arc<AtomicBool>,
release_first_token: Arc<Notify>,
server_stream_dropped: Arc<AtomicBool>,
}

Expand Down Expand Up @@ -120,6 +124,10 @@ impl pb::inference_server::Inference for FakeVllm {
"nested": {"flags": [true, null, "opaque"]},
});
let hang = self.hang.load(Ordering::SeqCst);
let hold_before_first_token = self.hold_before_first_token.load(Ordering::SeqCst);
let close_before_first_token = self.close_before_first_token.load(Ordering::SeqCst);
let first_token_pending = self.first_token_pending.clone();
let release_first_token = self.release_first_token.clone();
let dropped = self.server_stream_dropped.clone();

let stream = async_stream::try_stream! {
Expand Down Expand Up @@ -147,6 +155,15 @@ impl pb::inference_server::Inference for FakeVllm {
outputs: None,
};

if hold_before_first_token {
first_token_pending.store(true, Ordering::SeqCst);
release_first_token.notified().await;
first_token_pending.store(false, Ordering::SeqCst);
}
if close_before_first_token {
return;
}

if hang {
loop {
yield sequence_response(false, wants_logprobs, None);
Expand Down Expand Up @@ -482,6 +499,22 @@ fn request() -> PreprocessedRequest {
.expect("request")
}

fn decode_request() -> PreprocessedRequest {
let mut request = request();
request.prefill_result = Some(PrefillResult {
disaggregated_params: json!({
"do_remote_decode": false,
"do_remote_prefill": true,
"remote_engine_id": "prefill-0",
"remote_host": "127.0.0.1",
"remote_port": 20097,
"remote_block_ids": [7, 8],
}),
prompt_tokens_details: None,
});
request
}

fn engine(endpoint: &str, mode: DisaggregationMode, connections: usize) -> VllmSidecarEngine {
let transport = GrpcTransportConfig {
connections: NonZeroUsize::new(connections).expect("non-zero connection count"),
Expand Down Expand Up @@ -815,6 +848,108 @@ async fn cancellation_interrupts_pending_response_headers() {
server.service.release_headers.notify_waiters();
}

#[tokio::test]
async fn decode_cancellation_waits_for_submission_and_first_token() {
let service = FakeVllm::default();
service.hang_before_headers.store(true, Ordering::SeqCst);
service
.hold_before_first_token
.store(true, Ordering::SeqCst);
let server = FakeServer::start(service).await;
let engine = engine(&server.endpoint, DisaggregationMode::Decode, 1);
engine.start(0).await.expect("start");

let context = dynamo_backend_common::testing::mock_context();
let generate = engine.generate(
decode_request(),
GenerateContext::new(context.clone(), None),
);
tokio::pin!(generate);

tokio::select! {
_ = &mut generate => panic!("decode returned before response headers were gated"),
_ = async {
while !server.service.headers_pending.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
} => {}
}
assert_eq!(server.service.requests.lock().await.len(), 1);
context.stop_generating();
tokio::select! {
_ = &mut generate => panic!("decode cancellation returned before response headers"),
_ = tokio::time::sleep(std::time::Duration::from_millis(50)) => {}
}

server.service.release_headers.notify_one();
let mut stream = tokio::time::timeout(std::time::Duration::from_secs(2), &mut generate)
.await
.expect("decode response headers")
.expect("decode stream");
let next = stream.next();
tokio::pin!(next);
tokio::select! {
_ = &mut next => panic!("decode returned before the first token was gated"),
_ = async {
while !server.service.first_token_pending.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
} => {}
}
assert!(
!server.service.server_stream_dropped.load(Ordering::SeqCst),
"decode stream dropped before the first token"
);
tokio::select! {
_ = &mut next => panic!("decode cancellation completed before the first token"),
_ = tokio::time::sleep(std::time::Duration::from_millis(50)) => {}
}

server.service.release_first_token.notify_one();
let terminal = tokio::time::timeout(std::time::Duration::from_secs(2), &mut next)
.await
.expect("first token did not release decode cancellation")
.expect("cancelled terminal")
.expect("cancelled output");
assert_eq!(terminal.finish_reason, Some(FinishReason::Cancelled));
drop(stream);

tokio::time::timeout(std::time::Duration::from_secs(2), async {
while !server.service.server_stream_dropped.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
})
.await
.expect("server stream dropped after first token");
}

#[tokio::test]
async fn decode_cancellation_maps_premature_eof_to_cancelled() {
let service = FakeVllm::default();
service
.close_before_first_token
.store(true, Ordering::SeqCst);
let server = FakeServer::start(service).await;
let engine = engine(&server.endpoint, DisaggregationMode::Decode, 1);
engine.start(0).await.expect("start");

let context = dynamo_backend_common::testing::mock_context();
let mut stream = engine
.generate(
decode_request(),
GenerateContext::new(context.clone(), None),
)
.await
.expect("decode stream");
context.stop_generating();
let terminal = tokio::time::timeout(std::time::Duration::from_secs(2), stream.next())
.await
.expect("premature EOF did not release decode cancellation")
.expect("cancelled terminal")
.expect("cancelled output");
assert_eq!(terminal.finish_reason, Some(FinishReason::Cancelled));
}

#[tokio::test]
async fn unsupported_features_fail_before_rpc_submission() {
let server = FakeServer::start(FakeVllm::default()).await;
Expand Down
Loading