From cff83a4ae8fbcba9ceed23f0c4b176a41179cd1c Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 2 Aug 2026 15:02:54 +1000 Subject: [PATCH 1/5] fix(openai): cancel timed-out generation before releasing its lane --- crates/openai-frontend/src/backend.rs | 41 +++- .../openai-frontend/src/guardrails/compact.rs | 22 +- crates/openai-frontend/src/guardrails/mod.rs | 37 ++- crates/openai-frontend/src/hooks.rs | 24 +- crates/openai-frontend/src/router.rs | 221 ++++++++++++++++-- .../skippy-server/src/frontend/admission.rs | 21 +- crates/skippy-server/src/frontend/backend.rs | 82 ++++++- .../generation_flow/text_generation.rs | 9 +- .../src/frontend/linear_proposal/execution.rs | 19 ++ .../local_generation/linear_decode.rs | 13 ++ .../src/frontend/tests/generation.rs | 66 ++++++ 11 files changed, 517 insertions(+), 38 deletions(-) diff --git a/crates/openai-frontend/src/backend.rs b/crates/openai-frontend/src/backend.rs index 632919b5f..52bd662b7 100644 --- a/crates/openai-frontend/src/backend.rs +++ b/crates/openai-frontend/src/backend.rs @@ -8,6 +8,7 @@ use std::{ use async_trait::async_trait; use futures_core::Stream; +use tokio::sync::Notify; use crate::{ chat::{ChatCompletionChunk, ChatCompletionRequest, ChatCompletionResponse}, @@ -23,9 +24,15 @@ pub type CompletionStream = pub type OpenAiResult = Result; +#[derive(Debug, Default)] +struct CancellationState { + cancelled: AtomicBool, + notify: Notify, +} + #[derive(Debug, Clone, Default)] pub struct CancellationToken { - cancelled: Arc, + state: Arc, } impl CancellationToken { @@ -34,11 +41,23 @@ impl CancellationToken { } pub fn cancel(&self) { - self.cancelled.store(true, Ordering::Relaxed); + if !self.state.cancelled.swap(true, Ordering::AcqRel) { + self.state.notify.notify_waiters(); + } } pub fn is_cancelled(&self) -> bool { - self.cancelled.load(Ordering::Relaxed) + self.state.cancelled.load(Ordering::Acquire) + } + + pub async fn cancelled(&self) { + loop { + let notified = self.state.notify.notified(); + if self.is_cancelled() { + return; + } + notified.await; + } } } @@ -74,6 +93,14 @@ pub trait OpenAiBackend: Send + Sync + 'static { request: ChatCompletionRequest, ) -> OpenAiResult; + async fn chat_completion_with_context( + &self, + request: ChatCompletionRequest, + _context: OpenAiRequestContext, + ) -> OpenAiResult { + self.chat_completion(request).await + } + async fn chat_completion_stream( &self, request: ChatCompletionRequest, @@ -86,6 +113,14 @@ pub trait OpenAiBackend: Send + Sync + 'static { )) } + async fn completion_with_context( + &self, + request: CompletionRequest, + _context: OpenAiRequestContext, + ) -> OpenAiResult { + self.completion(request).await + } + async fn completion_stream( &self, _request: CompletionRequest, diff --git a/crates/openai-frontend/src/guardrails/compact.rs b/crates/openai-frontend/src/guardrails/compact.rs index 093b29940..57c8b337b 100644 --- a/crates/openai-frontend/src/guardrails/compact.rs +++ b/crates/openai-frontend/src/guardrails/compact.rs @@ -63,9 +63,18 @@ impl OpenAiBackend for CompactingOpenAiBackend { async fn chat_completion( &self, request: ChatCompletionRequest, + ) -> OpenAiResult { + self.chat_completion_with_context(request, OpenAiRequestContext::new()) + .await + } + + async fn chat_completion_with_context( + &self, + request: ChatCompletionRequest, + context: OpenAiRequestContext, ) -> OpenAiResult { self.backend - .chat_completion(self.compact_request(request)?) + .chat_completion_with_context(self.compact_request(request)?, context) .await } @@ -80,7 +89,16 @@ impl OpenAiBackend for CompactingOpenAiBackend { } async fn completion(&self, request: CompletionRequest) -> OpenAiResult { - self.backend.completion(request).await + self.completion_with_context(request, OpenAiRequestContext::new()) + .await + } + + async fn completion_with_context( + &self, + request: CompletionRequest, + context: OpenAiRequestContext, + ) -> OpenAiResult { + self.backend.completion_with_context(request, context).await } async fn completion_stream( diff --git a/crates/openai-frontend/src/guardrails/mod.rs b/crates/openai-frontend/src/guardrails/mod.rs index 0a050132e..09a4924fc 100644 --- a/crates/openai-frontend/src/guardrails/mod.rs +++ b/crates/openai-frontend/src/guardrails/mod.rs @@ -75,6 +75,7 @@ impl GuardedOpenAiBackend { async fn guarded_chat_completion( &self, request: ChatCompletionRequest, + context: OpenAiRequestContext, ) -> OpenAiResult { let _guardrail_error_catalog = guardrail_error_catalog(); let policy = self.policy.snapshot(); @@ -90,13 +91,15 @@ impl GuardedOpenAiBackend { GuardrailTelemetryOutcome::PassThrough, None, ); - self.backend.chat_completion(request).await + self.backend + .chat_completion_with_context(request, context) + .await } GuardrailRequestOutcome::Reject { kind } => Err(errors::guardrail_error(*kind)), GuardrailRequestOutcome::Guarded { backend_request } => { if matches!(policy.mode, GuardrailMode::MetricsOnly) { return self - .metrics_only_chat_completion(request, &engine, &prepared) + .metrics_only_chat_completion(request, &engine, &prepared, context) .await; } @@ -107,7 +110,7 @@ impl GuardedOpenAiBackend { loop { let response = self .backend - .chat_completion(attempt_request.clone()) + .chat_completion_with_context(attempt_request.clone(), context.clone()) .await?; let classified = engine.classify_response(&prepared, &response); let contract = telemetry_contract(&prepared.state.request_contract); @@ -165,8 +168,12 @@ impl GuardedOpenAiBackend { request: ChatCompletionRequest, engine: &GuardrailEngine, prepared: &state::PreparedGuardrailRequest, + context: OpenAiRequestContext, ) -> OpenAiResult { - let response = self.backend.chat_completion(request).await?; + let response = self + .backend + .chat_completion_with_context(request, context) + .await?; let classified = engine.classify_response(prepared, &response); self.record_outcome( prepared.state.mode, @@ -303,7 +310,16 @@ impl OpenAiBackend for GuardedOpenAiBackend { &self, request: ChatCompletionRequest, ) -> OpenAiResult { - self.guarded_chat_completion(request).await + self.guarded_chat_completion(request, OpenAiRequestContext::new()) + .await + } + + async fn chat_completion_with_context( + &self, + request: ChatCompletionRequest, + context: OpenAiRequestContext, + ) -> OpenAiResult { + self.guarded_chat_completion(request, context).await } async fn chat_completion_stream( @@ -315,7 +331,16 @@ impl OpenAiBackend for GuardedOpenAiBackend { } async fn completion(&self, request: CompletionRequest) -> OpenAiResult { - self.backend.completion(request).await + self.completion_with_context(request, OpenAiRequestContext::new()) + .await + } + + async fn completion_with_context( + &self, + request: CompletionRequest, + context: OpenAiRequestContext, + ) -> OpenAiResult { + self.backend.completion_with_context(request, context).await } async fn completion_stream( diff --git a/crates/openai-frontend/src/hooks.rs b/crates/openai-frontend/src/hooks.rs index 895149a53..b6192d057 100644 --- a/crates/openai-frontend/src/hooks.rs +++ b/crates/openai-frontend/src/hooks.rs @@ -128,12 +128,23 @@ impl OpenAiBackend for HookedOpenAiBackend { } async fn chat_completion( + &self, + request: ChatCompletionRequest, + ) -> OpenAiResult { + self.chat_completion_with_context(request, OpenAiRequestContext::new()) + .await + } + + async fn chat_completion_with_context( &self, mut request: ChatCompletionRequest, + context: OpenAiRequestContext, ) -> OpenAiResult { let outcome = self.hooks.before_chat_completion(&mut request).await?; apply_chat_hook_outcome(&mut request, &outcome); - self.backend.chat_completion(request).await + self.backend + .chat_completion_with_context(request, context) + .await } async fn chat_completion_stream( @@ -147,7 +158,16 @@ impl OpenAiBackend for HookedOpenAiBackend { } async fn completion(&self, request: CompletionRequest) -> OpenAiResult { - self.backend.completion(request).await + self.completion_with_context(request, OpenAiRequestContext::new()) + .await + } + + async fn completion_with_context( + &self, + request: CompletionRequest, + context: OpenAiRequestContext, + ) -> OpenAiResult { + self.backend.completion_with_context(request, context).await } async fn completion_stream( diff --git a/crates/openai-frontend/src/router.rs b/crates/openai-frontend/src/router.rs index f6e1e6ffc..21eb39712 100644 --- a/crates/openai-frontend/src/router.rs +++ b/crates/openai-frontend/src/router.rs @@ -169,10 +169,13 @@ async fn chat_completions( let model = request.model.clone(); let context = OpenAiRequestContext::new(); let cancellation = context.cancellation_token(); - let stream = backend_call( + let stream = backend_call_with_cancellation( &state, "chat_completion_stream", - state.backend.chat_completion_stream(request, context), + &context, + state + .backend + .chat_completion_stream(request, context.clone()), ) .await?; let prelude = stream::once(async move { json_event(&ChatCompletionChunk::role(model)) }); @@ -187,11 +190,15 @@ async fn chat_completions( .chain(stream::once(async { done_event() })); Ok(sse_response(events, cancellation)) } else { + let context = OpenAiRequestContext::new(); Ok(Json( - backend_call( + backend_call_with_cancellation( &state, "chat_completion", - state.backend.chat_completion(request), + &context, + state + .backend + .chat_completion_with_context(request, context.clone()), ) .await?, ) @@ -221,10 +228,13 @@ async fn responses( let context = OpenAiRequestContext::new(); let cancellation = context.cancellation_token(); let state_machine = Arc::new(Mutex::new(ResponseSseState::new(request.model.clone()))); - let stream = backend_call( + let stream = backend_call_with_cancellation( &state, "responses_stream", - state.backend.chat_completion_stream(request, context), + &context, + state + .backend + .chat_completion_stream(request, context.clone()), ) .await?; let body_state = state_machine.clone(); @@ -415,8 +425,16 @@ async fn responses( Ok(sse_response(events, cancellation)) } _ => { - let response = - backend_call(&state, "responses", state.backend.chat_completion(request)).await?; + let context = OpenAiRequestContext::new(); + let response = backend_call_with_cancellation( + &state, + "responses", + &context, + state + .backend + .chat_completion_with_context(request, context.clone()), + ) + .await?; let translated = translate_chat_completion_response_to_responses(&response)?; Ok(Json(translated).into_response()) } @@ -435,10 +453,11 @@ async fn completions( let include_usage = request.include_usage(); let context = OpenAiRequestContext::new(); let cancellation = context.cancellation_token(); - let stream = backend_call( + let stream = backend_call_with_cancellation( &state, "completion_stream", - state.backend.completion_stream(request, context), + &context, + state.backend.completion_stream(request, context.clone()), ) .await?; let events = stream @@ -452,10 +471,19 @@ async fn completions( .chain(stream::once(async { done_event() })); Ok(sse_response(events, cancellation)) } else { - Ok( - Json(backend_call(&state, "completion", state.backend.completion(request)).await?) - .into_response(), + let context = OpenAiRequestContext::new(); + Ok(Json( + backend_call_with_cancellation( + &state, + "completion", + &context, + state + .backend + .completion_with_context(request, context.clone()), + ) + .await?, ) + .into_response()) } } @@ -494,6 +522,59 @@ fn resolve_agent_session( } } +struct CancelOnDrop { + context: OpenAiRequestContext, + armed: bool, +} + +impl CancelOnDrop { + fn new(context: &OpenAiRequestContext) -> Self { + Self { + context: context.clone(), + armed: true, + } + } + + fn disarm(&mut self) { + self.armed = false; + } +} + +impl Drop for CancelOnDrop { + fn drop(&mut self) { + if self.armed { + self.context.cancel(); + } + } +} + +async fn backend_call_with_cancellation( + state: &FrontendState, + operation: &'static str, + context: &OpenAiRequestContext, + future: F, +) -> OpenAiResult +where + F: Future>, +{ + let mut cancel_on_drop = CancelOnDrop::new(context); + let result = match state.config.backend_timeout { + Some(timeout) => match tokio::time::timeout(timeout, future).await { + Ok(result) => result, + Err(_) => { + context.cancel(); + return Err(OpenAiError::timeout(format!( + "{operation} timed out after {} ms", + timeout.as_millis() + ))); + } + }, + None => future.await, + }; + cancel_on_drop.disarm(); + result +} + async fn backend_call( state: &FrontendState, operation: &'static str, @@ -867,6 +948,42 @@ mod tests { unreachable!("agent-session tests use non-streaming requests") } } + + struct NonStreamingCancellationBackend { + token: Arc>>, + } + + #[async_trait] + impl OpenAiBackend for NonStreamingCancellationBackend { + async fn models(&self) -> OpenAiResult> { + Ok(vec![ModelObject::new("cancel-model")]) + } + + async fn chat_completion( + &self, + _request: ChatCompletionRequest, + ) -> OpenAiResult { + unreachable!("frontend must use the context-aware non-streaming path") + } + + async fn chat_completion_with_context( + &self, + _request: ChatCompletionRequest, + context: OpenAiRequestContext, + ) -> OpenAiResult { + *self.token.lock().expect("token lock poisoned") = Some(context.cancellation_token()); + std::future::pending().await + } + + async fn chat_completion_stream( + &self, + _request: ChatCompletionRequest, + _context: OpenAiRequestContext, + ) -> OpenAiResult { + unreachable!("cancellation backend test only calls non-streaming") + } + } + #[async_trait] impl OpenAiBackend for CancellationBackend { async fn models(&self) -> OpenAiResult> { @@ -1390,6 +1507,84 @@ mod tests { assert!(cancellation.is_cancelled()); } + #[tokio::test] + async fn non_streaming_timeout_cancels_request_context() { + let token = Arc::new(Mutex::new(None)); + let app = router_for_with_config( + Arc::new(NonStreamingCancellationBackend { + token: token.clone(), + }), + OpenAiFrontendConfig::default().with_backend_timeout(Duration::from_millis(5)), + ); + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header("content-type", "application/json") + .body(Body::from( + json!({ + "model": "cancel-model", + "messages": [{"role": "user", "content": "hello"}] + }) + .to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::GATEWAY_TIMEOUT); + let cancellation = token + .lock() + .expect("token lock poisoned") + .clone() + .expect("backend saw request context"); + assert!(cancellation.is_cancelled()); + } + + #[tokio::test] + async fn dropping_non_streaming_request_cancels_request_context() { + let token = Arc::new(Mutex::new(None)); + let app = router_for_with_config( + Arc::new(NonStreamingCancellationBackend { + token: token.clone(), + }), + OpenAiFrontendConfig::default().without_backend_timeout(), + ); + let request = Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header("content-type", "application/json") + .body(Body::from( + json!({ + "model": "cancel-model", + "messages": [{"role": "user", "content": "hello"}] + }) + .to_string(), + )) + .unwrap(); + let task = tokio::spawn(app.oneshot(request)); + + let cancellation = tokio::time::timeout(Duration::from_millis(100), async { + loop { + if let Some(token) = token.lock().expect("token lock poisoned").clone() { + break token; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("backend received request context"); + assert!(!cancellation.is_cancelled()); + + task.abort(); + let _ = task.await; + tokio::time::timeout(Duration::from_millis(100), cancellation.cancelled()) + .await + .expect("dropped request cancelled backend context"); + } + #[tokio::test] async fn chat_completion_route_maps_backend_errors() { let response = post_json( diff --git a/crates/skippy-server/src/frontend/admission.rs b/crates/skippy-server/src/frontend/admission.rs index 17bc4015e..ab4fd7a7d 100644 --- a/crates/skippy-server/src/frontend/admission.rs +++ b/crates/skippy-server/src/frontend/admission.rs @@ -45,10 +45,20 @@ impl GenerationTokenBudget { } } + #[cfg(test)] pub(super) fn reserve( self: &Arc, request: GenerationTokenBudgetRequest, admission_timeout: Duration, + ) -> OpenAiResult { + self.reserve_cancellable(request, admission_timeout, None) + } + + pub(super) fn reserve_cancellable( + self: &Arc, + request: GenerationTokenBudgetRequest, + admission_timeout: Duration, + cancellation: Option<&openai_frontend::CancellationToken>, ) -> OpenAiResult { let tokens = request.reservation_tokens(self.capacity_tokens); let deadline = Instant::now() + admission_timeout; @@ -57,6 +67,9 @@ impl GenerationTokenBudget { .lock() .map_err(|_| OpenAiError::backend("generation token budget lock poisoned"))?; loop { + if cancellation.is_some_and(openai_frontend::CancellationToken::is_cancelled) { + return Err(OpenAiError::backend("request cancelled")); + } if state.active_tokens.saturating_add(tokens) <= self.capacity_tokens { state.active_tokens = state.active_tokens.saturating_add(tokens); return Ok(GenerationTokenReservation { @@ -76,13 +89,17 @@ impl GenerationTokenBudget { )); } - let wait_for = deadline.saturating_duration_since(now); + let wait_for = deadline.saturating_duration_since(now).min( + cancellation + .map(|_| Duration::from_millis(10)) + .unwrap_or(admission_timeout), + ); let (next_state, wait_result) = self .released .wait_timeout(state, wait_for) .map_err(|_| OpenAiError::backend("generation token budget lock poisoned"))?; state = next_state; - if wait_result.timed_out() { + if wait_result.timed_out() && Instant::now() >= deadline { return Err(generation_token_budget_timeout_error( admission_timeout, tokens, diff --git a/crates/skippy-server/src/frontend/backend.rs b/crates/skippy-server/src/frontend/backend.rs index 7ddec6402..6d77ee2cb 100644 --- a/crates/skippy-server/src/frontend/backend.rs +++ b/crates/skippy-server/src/frontend/backend.rs @@ -51,6 +51,26 @@ use tokio::sync::OwnedSemaphorePermit; use tokio::sync::mpsc; use tokio::task; +fn request_cancelled_error() -> OpenAiError { + OpenAiError::backend("request cancelled") +} + +pub(in crate::frontend) async fn run_blocking_generation_worker( + permit: OwnedSemaphorePermit, + context: OpenAiRequestContext, + work: F, +) -> Result +where + T: Send + 'static, + F: FnOnce(openai_frontend::CancellationToken) -> T + Send + 'static, +{ + task::spawn_blocking(move || { + let _permit = permit; + work(context.cancellation_token()) + }) + .await +} + #[async_trait] impl OpenAiBackend for StageOpenAiBackend { async fn models(&self) -> OpenAiResult> { @@ -58,8 +78,17 @@ impl OpenAiBackend for StageOpenAiBackend { } async fn chat_completion( + &self, + request: ChatCompletionRequest, + ) -> OpenAiResult { + self.chat_completion_with_context(request, OpenAiRequestContext::new()) + .await + } + + async fn chat_completion_with_context( &self, mut request: ChatCompletionRequest, + context: OpenAiRequestContext, ) -> OpenAiResult { let ids = OpenAiGenerationIds::new( OpenAiCacheHints::from_chat_request(&request), @@ -105,6 +134,7 @@ impl OpenAiBackend for StageOpenAiBackend { request.stop.clone(), sampling, Some(request.clone()), + context, ids.clone(), ) .await?; @@ -220,7 +250,16 @@ impl OpenAiBackend for StageOpenAiBackend { }))) } - async fn completion(&self, mut request: CompletionRequest) -> OpenAiResult { + async fn completion(&self, request: CompletionRequest) -> OpenAiResult { + self.completion_with_context(request, OpenAiRequestContext::new()) + .await + } + + async fn completion_with_context( + &self, + mut request: CompletionRequest, + context: OpenAiRequestContext, + ) -> OpenAiResult { let ids = OpenAiGenerationIds::new( OpenAiCacheHints::from_completion_request(&request), request.agent_session(), @@ -251,6 +290,7 @@ impl OpenAiBackend for StageOpenAiBackend { request.stop.clone(), sampling, None, + context, ids.clone(), ) .await?; @@ -448,6 +488,7 @@ impl StageOpenAiBackend { Ok(()) } + #[allow(clippy::too_many_arguments)] async fn run_generation( &self, prompt: PreparedGenerationPrompt, @@ -455,10 +496,17 @@ impl StageOpenAiBackend { stop: Option, sampling: SamplingConfig, hook_request: Option, + context: OpenAiRequestContext, ids: OpenAiGenerationIds, ) -> OpenAiResult { let admit_timer = PhaseTimer::start(); - let permit = self.acquire_generation_permit().await?; + let cancellation = context.cancellation_token(); + let permit = tokio::select! { + permit = self.acquire_generation_permit() => permit?, + () = cancellation.cancelled() => { + return Err(request_cancelled_error()); + } + }; let mut admit_attrs = self.openai_attrs(&ids); admit_attrs.insert( "llama_stage.openai_phase".to_string(), @@ -467,22 +515,32 @@ impl StageOpenAiBackend { self.emit_openai_phase("stage.openai_generation_admit", admit_timer, admit_attrs); let backend = self.clone(); let hook_runtime = Some(tokio::runtime::Handle::current()); - task::spawn_blocking(move || { - let _permit = permit; - backend.generate_text( + let worker_context = context.clone(); + let result = run_blocking_generation_worker(permit, worker_context.clone(), move |token| { + let output = backend.generate_text( prompt, max_tokens, stop.as_ref(), sampling, hook_request, hook_runtime, - None, + Some(&token), ids, |_| Ok(()), - ) + ); + if worker_context.is_cancelled() { + Err(request_cancelled_error()) + } else { + output + } }) .await - .map_err(|error| OpenAiError::backend(format!("generation task failed: {error}")))? + .map_err(|error| OpenAiError::backend(format!("generation task failed: {error}")))?; + if context.is_cancelled() { + Err(request_cancelled_error()) + } else { + result + } } #[allow(clippy::too_many_arguments)] @@ -500,7 +558,13 @@ impl StageOpenAiBackend { ids: OpenAiGenerationIds, ) -> OpenAiResult { let admit_timer = PhaseTimer::start(); - let permit = self.acquire_generation_permit().await?; + let cancellation = context.cancellation_token(); + let permit = tokio::select! { + permit = self.acquire_generation_permit() => permit?, + () = cancellation.cancelled() => { + return Err(request_cancelled_error()); + } + }; let mut admit_attrs = self.openai_attrs(&ids); admit_attrs.insert( "llama_stage.openai_phase".to_string(), diff --git a/crates/skippy-server/src/frontend/generation_flow/text_generation.rs b/crates/skippy-server/src/frontend/generation_flow/text_generation.rs index e211c7d8c..67a7c9fd8 100644 --- a/crates/skippy-server/src/frontend/generation_flow/text_generation.rs +++ b/crates/skippy-server/src/frontend/generation_flow/text_generation.rs @@ -24,6 +24,9 @@ impl StageOpenAiBackend { on_text_chunk: impl FnMut(&str) -> OpenAiResult<()>, ) -> OpenAiResult { let generation_timer = PhaseTimer::start(); + if cancellation.is_some_and(openai_frontend::CancellationToken::is_cancelled) { + return Err(OpenAiError::backend("request cancelled")); + } if prompt.text.is_empty() { return Err(OpenAiError::invalid_request( "request prompt/messages produced no text", @@ -50,6 +53,9 @@ impl StageOpenAiBackend { .collect::>(); let tokenize_timer = PhaseTimer::start(); let prompt_token_ids = self.tokenize(&prompt.text)?; + if cancellation.is_some_and(openai_frontend::CancellationToken::is_cancelled) { + return Err(OpenAiError::backend("request cancelled")); + } let mut tokenize_attrs = self.openai_attrs(&ids); tokenize_attrs.insert( "llama_stage.prompt_chars".to_string(), @@ -65,9 +71,10 @@ impl StageOpenAiBackend { } let max_tokens = max_tokens.resolve(prompt_token_ids.len(), self.ctx_size)?; let token_admit_timer = PhaseTimer::start(); - let token_budget_reservation = self.generation_token_budget.reserve( + let token_budget_reservation = self.generation_token_budget.reserve_cancellable( GenerationTokenBudgetRequest::new(prompt_token_ids.len(), max_tokens), GENERATION_ADMISSION_TIMEOUT, + cancellation, )?; let mut token_admit_attrs = self.openai_attrs(&ids); token_admit_attrs.insert( diff --git a/crates/skippy-server/src/frontend/linear_proposal/execution.rs b/crates/skippy-server/src/frontend/linear_proposal/execution.rs index 61480cb6e..a9abc02a5 100644 --- a/crates/skippy-server/src/frontend/linear_proposal/execution.rs +++ b/crates/skippy-server/src/frontend/linear_proposal/execution.rs @@ -60,8 +60,10 @@ impl StageOpenAiBackend { &self, params: LinearProposalExecutionParams<'_>, queried: QueriedLinearProposal, + cancellation: Option<&openai_frontend::CancellationToken>, on_token: &mut impl FnMut(i32) -> OpenAiResult, ) -> OpenAiResult> { + ensure_request_active(cancellation)?; let proposal_token_count = queried.proposal.token_ids.len(); let mut verify_inputs = Vec::with_capacity(proposal_token_count.saturating_add(1)); verify_inputs.push(params.current); @@ -71,6 +73,7 @@ impl StageOpenAiBackend { params, &queried.proposal.token_ids, &verify_inputs, + cancellation, on_token, )? else { @@ -134,8 +137,10 @@ impl StageOpenAiBackend { params: LinearProposalExecutionParams<'_>, proposal_tokens: &[i32], verify_inputs: &[i32], + cancellation: Option<&openai_frontend::CancellationToken>, on_token: &mut impl FnMut(i32) -> OpenAiResult, ) -> OpenAiResult> { + ensure_request_active(cancellation)?; let verify_timer = Instant::now(); let verify_lock_timer = Instant::now(); let mut runtime = self @@ -192,6 +197,10 @@ impl StageOpenAiBackend { let mut reached_stop = false; let mut callback_error = None; for token in predictions.iter().copied().take(decision.commit_count) { + if cancellation.is_some_and(openai_frontend::CancellationToken::is_cancelled) { + callback_error = Some(OpenAiError::backend("request cancelled")); + break; + } committed_tokens.push(token); match on_token(token) { Ok(TokenControl::Continue) => {} @@ -323,6 +332,16 @@ impl StageOpenAiBackend { } } +fn ensure_request_active( + cancellation: Option<&openai_frontend::CancellationToken>, +) -> OpenAiResult<()> { + if cancellation.is_some_and(openai_frontend::CancellationToken::is_cancelled) { + Err(OpenAiError::backend("request cancelled")) + } else { + Ok(()) + } +} + pub(crate) fn elapsed_us(started: Instant) -> u64 { u64::try_from(started.elapsed().as_micros()).unwrap_or(u64::MAX) } diff --git a/crates/skippy-server/src/frontend/local_generation/linear_decode.rs b/crates/skippy-server/src/frontend/local_generation/linear_decode.rs index 0fe181fff..0492524c6 100644 --- a/crates/skippy-server/src/frontend/local_generation/linear_decode.rs +++ b/crates/skippy-server/src/frontend/local_generation/linear_decode.rs @@ -28,6 +28,12 @@ impl StageOpenAiBackend { let Some(committed_token_ids) = state.linear_context_tokens.as_mut() else { return Ok(LinearProposalProgress::NotUsed); }; + if request + .cancellation + .is_some_and(openai_frontend::CancellationToken::is_cancelled) + { + return Err(OpenAiError::backend("request cancelled")); + } let remaining_new_tokens = (request.max_tokens as usize).saturating_sub(state.decoded_tokens); // Prefill leaves the final prompt token undecoded. When whole-prompt @@ -81,6 +87,12 @@ impl StageOpenAiBackend { let Some(queried) = queried else { return Ok(LinearProposalProgress::NotUsed); }; + if request + .cancellation + .is_some_and(openai_frontend::CancellationToken::is_cancelled) + { + return Err(OpenAiError::backend("request cancelled")); + } let decision_id = queried.proposal.decision_id.clone(); let receipt = execute_linear_proposal_with_terminal_discard(config, &decision_id, || { self.execute_local_linear_proposal( @@ -95,6 +107,7 @@ impl StageOpenAiBackend { prompt_token_count: request.prompt_token_ids.len(), }, queried, + request.cancellation, emit_token, ) })?; diff --git a/crates/skippy-server/src/frontend/tests/generation.rs b/crates/skippy-server/src/frontend/tests/generation.rs index 872de105f..ec316c999 100644 --- a/crates/skippy-server/src/frontend/tests/generation.rs +++ b/crates/skippy-server/src/frontend/tests/generation.rs @@ -1,4 +1,6 @@ use super::*; +use crate::frontend::backend::run_blocking_generation_worker; +use std::sync::atomic::AtomicBool; fn assert_generation_rate_limit(error: OpenAiError, message_fragment: &str) { assert_eq!(error.status(), StatusCode::TOO_MANY_REQUESTS); @@ -100,6 +102,70 @@ async fn generation_admission_waits_for_released_lane() { drop(permit); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn cancelled_worker_stops_before_the_next_request_acquires_the_only_lane() { + let generation_limit = Arc::new(Semaphore::new(1)); + let queue_depth = Arc::new(AtomicUsize::new(0)); + let first_permit = acquire_generation_permit_with_queue( + generation_limit.clone(), + queue_depth.clone(), + 1, + Duration::from_secs(1), + ) + .await + .unwrap(); + let first_context = OpenAiRequestContext::new(); + let worker_started = Arc::new(AtomicBool::new(false)); + let worker_stopped = Arc::new(AtomicBool::new(false)); + let started = worker_started.clone(); + let stopped = worker_stopped.clone(); + let first_worker = tokio::spawn(run_blocking_generation_worker( + first_permit, + first_context.clone(), + move |cancellation| { + started.store(true, Ordering::Release); + while !cancellation.is_cancelled() { + std::thread::sleep(Duration::from_millis(1)); + } + stopped.store(true, Ordering::Release); + }, + )); + + tokio::time::timeout(Duration::from_millis(100), async { + while !worker_started.load(Ordering::Acquire) { + tokio::task::yield_now().await; + } + }) + .await + .expect("first generation worker started"); + + let second_limit = generation_limit.clone(); + let second_queue_depth = queue_depth.clone(); + let second = tokio::spawn(async move { + let permit = acquire_generation_permit_with_queue( + second_limit, + second_queue_depth, + 1, + Duration::from_secs(1), + ) + .await + .unwrap(); + assert!( + worker_stopped.load(Ordering::Acquire), + "the first worker must stop before its permit is released" + ); + drop(permit); + }); + + first_context.cancel(); + tokio::time::timeout(Duration::from_millis(200), second) + .await + .expect("second request acquired the lane promptly") + .unwrap(); + first_worker.await.unwrap().unwrap(); + assert_eq!(generation_limit.available_permits(), 1); +} + #[test] fn trims_at_first_stop_sequence() { assert_eq!(trim_at_stop("hello END world", &["END"]), "hello "); From c759cf0d140ed0aa9754fd35739ca6e75683c303 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 2 Aug 2026 20:22:36 +1000 Subject: [PATCH 2/5] fix(skippy): admit fresh sessions before batch query --- crates/skippy-server/src/frontend/local_generation/tests.rs | 5 ++++- crates/skippy-server/src/runtime_state/frame_operations.rs | 2 +- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/crates/skippy-server/src/frontend/local_generation/tests.rs b/crates/skippy-server/src/frontend/local_generation/tests.rs index 18b323f80..b7523d4a1 100644 --- a/crates/skippy-server/src/frontend/local_generation/tests.rs +++ b/crates/skippy-server/src/frontend/local_generation/tests.rs @@ -163,7 +163,10 @@ fn local_generation_eventually_delivers_receipts_and_cleanup_survives_sink_error decode_frame_batcher, }; let sampling = SamplingConfig::default(); - let prompt_token_ids = [1]; + // A multi-token prompt takes the whole-prompt prefill path. Keep this + // above one token so the test exercises a fresh runtime session before + // its batch size is queried. + let prompt_token_ids = [1, 2]; let ids = OpenAiGenerationIds::new(OpenAiCacheHints::default(), None); let mut emitted = Vec::new(); backend.generate_local_tokens( diff --git a/crates/skippy-server/src/runtime_state/frame_operations.rs b/crates/skippy-server/src/runtime_state/frame_operations.rs index 2a3693f53..e47a21962 100644 --- a/crates/skippy-server/src/runtime_state/frame_operations.rs +++ b/crates/skippy-server/src/runtime_state/frame_operations.rs @@ -115,7 +115,7 @@ impl RuntimeState { } pub fn session_batch_size(&mut self, session_id: &str) -> Result { - self.active_session(session_id)?.batch_size() + self.session(session_id)?.batch_size() } pub fn ensure_session_active(&mut self, session_id: &str) -> Result<()> { From 9a3703fc4019132af43a1674ff53d3d1ea2e6bb8 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Tue, 4 Aug 2026 11:10:53 +1000 Subject: [PATCH 3/5] fix(skippy): distinguish session batch-size lookups --- .../src/frontend/local_generation/token_generation.rs | 4 ++-- .../skippy-server/src/kv_integration/resident_prefix.rs | 2 +- .../skippy-server/src/runtime_state/frame_operations.rs | 8 +++++++- 3 files changed, 10 insertions(+), 4 deletions(-) diff --git a/crates/skippy-server/src/frontend/local_generation/token_generation.rs b/crates/skippy-server/src/frontend/local_generation/token_generation.rs index e33610662..af115f14c 100644 --- a/crates/skippy-server/src/frontend/local_generation/token_generation.rs +++ b/crates/skippy-server/src/frontend/local_generation/token_generation.rs @@ -184,7 +184,7 @@ impl StageOpenAiBackend { .ensure_session_active(session_id) .map_err(openai_backend_error)?; let batch_size = runtime - .session_batch_size(session_id) + .admit_session_batch_size(session_id) .map_err(openai_backend_error)?; Ok(prompt_fits_single_prefill_sample( request.prompt_token_ids.len(), @@ -760,7 +760,7 @@ impl StageOpenAiBackend { .lock() .map_err(|_| OpenAiError::backend("runtime lock poisoned"))?; runtime - .session_batch_size(session_id) + .admit_session_batch_size(session_id) .map_err(openai_backend_error)? .saturating_sub(1) } else { diff --git a/crates/skippy-server/src/kv_integration/resident_prefix.rs b/crates/skippy-server/src/kv_integration/resident_prefix.rs index e17d9ee4b..7534e2992 100644 --- a/crates/skippy-server/src/kv_integration/resident_prefix.rs +++ b/crates/skippy-server/src/kv_integration/resident_prefix.rs @@ -186,7 +186,7 @@ impl KvStageIntegration { runtime: &mut RuntimeState, session_id: &str, ) -> Result { - let target_tokens = runtime.session_batch_size(session_id)? as u64; + let target_tokens = runtime.active_session_batch_size(session_id)? as u64; self.evict_resident_prefix_for_tokens(runtime, session_id, target_tokens) } diff --git a/crates/skippy-server/src/runtime_state/frame_operations.rs b/crates/skippy-server/src/runtime_state/frame_operations.rs index e47a21962..75c04d220 100644 --- a/crates/skippy-server/src/runtime_state/frame_operations.rs +++ b/crates/skippy-server/src/runtime_state/frame_operations.rs @@ -114,10 +114,16 @@ impl RuntimeState { result } - pub fn session_batch_size(&mut self, session_id: &str) -> Result { + /// Returns a session's batch size, admitting a new session when necessary. + pub fn admit_session_batch_size(&mut self, session_id: &str) -> Result { self.session(session_id)?.batch_size() } + /// Returns the batch size of an already admitted session. + pub fn active_session_batch_size(&self, session_id: &str) -> Result { + self.active_session(session_id)?.batch_size() + } + pub fn ensure_session_active(&mut self, session_id: &str) -> Result<()> { self.session(session_id).map(|_| ()) } From b5369930f9c86beb67c25ad6e3e6df11ccbe21d1 Mon Sep 17 00:00:00 2001 From: Nick DiZazzo Date: Tue, 4 Aug 2026 14:46:21 -0400 Subject: [PATCH 4/5] fix(skippy): accept mutable active batch lookup --- crates/skippy-server/src/runtime_state/frame_operations.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/skippy-server/src/runtime_state/frame_operations.rs b/crates/skippy-server/src/runtime_state/frame_operations.rs index 75c04d220..24111115e 100644 --- a/crates/skippy-server/src/runtime_state/frame_operations.rs +++ b/crates/skippy-server/src/runtime_state/frame_operations.rs @@ -120,7 +120,7 @@ impl RuntimeState { } /// Returns the batch size of an already admitted session. - pub fn active_session_batch_size(&self, session_id: &str) -> Result { + pub fn active_session_batch_size(&mut self, session_id: &str) -> Result { self.active_session(session_id)?.batch_size() } From b95c08d2d6de14c2a5fcb4c7aa16d63ba9c89a13 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Wed, 5 Aug 2026 10:50:14 +1000 Subject: [PATCH 5/5] fix(skippy): stop cancelled generation before token output --- .../generation_flow/text_generation.rs | 3 ++ .../src/frontend/linear_proposal/execution.rs | 32 +++++++++++++++---- .../local_generation/linear_decode.rs | 6 ++-- .../local_generation/token_generation.rs | 12 +++++++ 4 files changed, 44 insertions(+), 9 deletions(-) diff --git a/crates/skippy-server/src/frontend/generation_flow/text_generation.rs b/crates/skippy-server/src/frontend/generation_flow/text_generation.rs index 67a7c9fd8..6f6520ee7 100644 --- a/crates/skippy-server/src/frontend/generation_flow/text_generation.rs +++ b/crates/skippy-server/src/frontend/generation_flow/text_generation.rs @@ -76,6 +76,9 @@ impl StageOpenAiBackend { GENERATION_ADMISSION_TIMEOUT, cancellation, )?; + if cancellation.is_some_and(openai_frontend::CancellationToken::is_cancelled) { + return Err(OpenAiError::backend("request cancelled")); + } let mut token_admit_attrs = self.openai_attrs(&ids); token_admit_attrs.insert( "llama_stage.prompt_token_count".to_string(), diff --git a/crates/skippy-server/src/frontend/linear_proposal/execution.rs b/crates/skippy-server/src/frontend/linear_proposal/execution.rs index a9abc02a5..fccc66bd0 100644 --- a/crates/skippy-server/src/frontend/linear_proposal/execution.rs +++ b/crates/skippy-server/src/frontend/linear_proposal/execution.rs @@ -214,12 +214,6 @@ impl StageOpenAiBackend { } } } - if committed_tokens.is_empty() { - return Err(OpenAiError::backend( - "linear proposal classifier committed no target prediction", - )); - } - let canonical_position = params .base_position .checked_add( @@ -240,6 +234,12 @@ impl StageOpenAiBackend { }) })?; + if committed_tokens.is_empty() { + return Err(OpenAiError::backend( + "linear proposal classifier committed no target prediction", + )); + } + Ok(Some(LinearProposalExecution { decision, predictions, @@ -425,4 +425,24 @@ mod tests { .contains("synthetic callback failure") ); } + + #[test] + fn cancellation_error_is_returned_only_after_repair_runs() { + let repair_ran = Cell::new(false); + let result = finish_linear_proposal_after_repair( + Some(OpenAiError::backend("request cancelled")), + || { + repair_ran.set(true); + Ok(()) + }, + ); + + assert!(repair_ran.get()); + assert!( + result + .unwrap_err() + .to_string() + .contains("request cancelled") + ); + } } diff --git a/crates/skippy-server/src/frontend/local_generation/linear_decode.rs b/crates/skippy-server/src/frontend/local_generation/linear_decode.rs index 0492524c6..8e21123d0 100644 --- a/crates/skippy-server/src/frontend/local_generation/linear_decode.rs +++ b/crates/skippy-server/src/frontend/local_generation/linear_decode.rs @@ -84,15 +84,15 @@ impl StageOpenAiBackend { } LinearProposalQueryOutcome::Ready(queried) => Some(queried), }; - let Some(queried) = queried else { - return Ok(LinearProposalProgress::NotUsed); - }; if request .cancellation .is_some_and(openai_frontend::CancellationToken::is_cancelled) { return Err(OpenAiError::backend("request cancelled")); } + let Some(queried) = queried else { + return Ok(LinearProposalProgress::NotUsed); + }; let decision_id = queried.proposal.decision_id.clone(); let receipt = execute_linear_proposal_with_terminal_discard(config, &decision_id, || { self.execute_local_linear_proposal( diff --git a/crates/skippy-server/src/frontend/local_generation/token_generation.rs b/crates/skippy-server/src/frontend/local_generation/token_generation.rs index af115f14c..4bd95073c 100644 --- a/crates/skippy-server/src/frontend/local_generation/token_generation.rs +++ b/crates/skippy-server/src/frontend/local_generation/token_generation.rs @@ -125,6 +125,12 @@ impl StageOpenAiBackend { Ok(control) }; let result = (|| { + if request + .cancellation + .is_some_and(openai_frontend::CancellationToken::is_cancelled) + { + return Err(OpenAiError::backend("request cancelled")); + } let prefill = self.prefill_prompt(&request, &session_id, &mut cache_stats)?; self.configure_chat_sampling_if_needed( &request, @@ -742,6 +748,12 @@ impl StageOpenAiBackend { .expect("checked non-empty prompt"); let mut stopped = false; if let Some(predicted) = prompt_prefill_sample { + if request + .cancellation + .is_some_and(openai_frontend::CancellationToken::is_cancelled) + { + return Err(OpenAiError::backend("request cancelled")); + } current = predicted; decoded_tokens += 1; stopped = emit_token(current)? == TokenControl::Stop;