diff --git a/crates/tokenizer/src/stop.rs b/crates/tokenizer/src/stop.rs index 3b0ed3e6fa..4f9e7cc1f7 100644 --- a/crates/tokenizer/src/stop.rs +++ b/crates/tokenizer/src/stop.rs @@ -376,6 +376,8 @@ mod tests { // Process stop token let result = decoder.process_token(999).unwrap(); // assert_eq!(result, SequenceDecoderOutput::Stopped); + // Token-level stops don't record a matched string stop. + assert_eq!(decoder.matched_stop(), None); // Further tokens should also return Stopped let result = decoder.process_token(2).unwrap(); @@ -620,6 +622,7 @@ mod tests { decoder.is_stopped(), "Decoder should be stopped after the full stop sequence match" ); + assert_eq!(decoder.matched_stop(), Some("Hello world")); // Any further tokens should also return Stopped let result3 = decoder.process_token(3).unwrap(); diff --git a/model_gateway/src/routers/grpc/regular/streaming.rs b/model_gateway/src/routers/grpc/regular/streaming.rs index 71427e68bc..507ea537de 100644 --- a/model_gateway/src/routers/grpc/regular/streaming.rs +++ b/model_gateway/src/routers/grpc/regular/streaming.rs @@ -292,6 +292,10 @@ impl StreamingProcessor { let mut completion_tokens = CompletionTokenTracker::new(); let mut cached_tokens: HashMap = HashMap::new(); let mut reasoning_tokens: HashMap = HashMap::new(); + let mut terminal_indices: HashSet = HashSet::new(); + let expected_choices = original_request.n.unwrap_or(1).max(1) as usize; + let mut has_router_stop = false; + let mut router_terminated = false; // Parser state (lazy initialization per index) type PooledReasoningParser = Arc>>; @@ -453,7 +457,26 @@ impl StreamingProcessor { stopped_indices.insert(index); } + // A router-matched string stop is terminal for the public + // request but not for the backend (it never saw the string). + // Track it so the stream can be aborted — rather than drained — + // once every choice is terminal. + let router_string_stop = should_stop && stop_decoder.matched_stop().is_some(); + if router_string_stop { + has_router_stop = true; + prompt_tokens.insert(index, chunk.prompt_tokens()); + cached_tokens.insert(index, chunk.cached_tokens()); + reasoning_tokens.insert(index, chunk.reasoning_tokens()); + terminal_indices.insert(index); + } + let all_terminal = + has_router_stop && terminal_indices.len() >= expected_choices; + if chunk_text.is_empty() { + if all_terminal { + router_terminated = true; + break; + } continue; } @@ -515,6 +538,10 @@ impl StreamingProcessor { let tool_choice_enabled = !matches!(tool_choice, Some(ToolChoice::Value(ToolChoiceValue::None))); + let tool_parser_active = tools.is_some() + && !in_reasoning + && tool_choice_enabled + && (tool_parser_available || used_json_schema); if let Some(tools_ref) = tools.as_ref() { if !in_reasoning && tool_choice_enabled @@ -558,15 +585,11 @@ impl StreamingProcessor { .await .map_err(|_| "Failed to send tool call chunk".to_string())?; } - - // Always skip regular content when tool parsing is active - // Parser either emitted chunks or buffered content - continue; } } // Regular content emission - if !delta.is_empty() { + if !tool_parser_active && !delta.is_empty() { let content_chunk = ChatCompletionStreamResponse::builder(request_id, model) .created(created) @@ -583,6 +606,11 @@ impl StreamingProcessor { .await .map_err(|_| "Failed to send content chunk".to_string())?; } + + if all_terminal { + router_terminated = true; + break; + } } ProtoResponseVariant::Complete(complete) => { let index = complete.index(); @@ -626,8 +654,12 @@ impl StreamingProcessor { finish_reasons.insert(index, complete.finish_reason().to_string()); matched_stops.insert(index, complete.matched_stop_json()); } + terminal_indices.insert(index); - // Don't break - continue reading all Complete messages for n>1 + if has_router_stop && terminal_indices.len() >= expected_choices { + router_terminated = true; + break; + } } ProtoResponseVariant::None => continue, } @@ -725,8 +757,12 @@ impl StreamingProcessor { } } - // Mark stream as completed successfully to prevent abort on drop - grpc_stream.mark_completed(); + // A router-side string stop is terminal for the public request but not + // for the backend. Leave the stream unmarked so Drop sends its exact-ID + // Abort RPC instead of silently draining post-stop generation. + if !router_terminated { + grpc_stream.mark_completed(); + } // Record streaming metrics let total_prompt: u32 = prompt_tokens.values().copied().max().unwrap_or(0); @@ -1904,6 +1940,9 @@ impl StreamingProcessor { // Set once the local stop decoder fires: pins "stop" and ignores later // engine output (the backend has no stop-string detection over ZMQ). let mut stopped = false; + // Set when a router-matched string stop terminates the request: the + // backend stream is left unmarked so Drop aborts it instead of draining. + let mut router_terminated = false; // Per-request effective parser names (model-card override → configured). let reasoning_parser_name = self.parser_resolver.reasoning_parser(model); @@ -2059,12 +2098,23 @@ impl StreamingProcessor { .map(|s| Value::String(s.to_string())); } + // A router-matched string stop is terminal for the public + // request but not for the backend (it never saw the string): + // emit any pre-stop text below, then break and leave the + // stream to abort on drop instead of draining it. + let router_string_stop = should_stop && matched_stop.is_some(); + if chunk_text.is_empty() { + if router_string_stop { + prompt_tokens = chunk.prompt_tokens(); + router_terminated = true; + break; + } continue; } // Apply reasoning parser - let (normal_text, reasoning_chunk_text, in_reasoning) = + let (mut normal_text, reasoning_chunk_text, in_reasoning) = if reasoning_parser_available { self.process_messages_reasoning( &chunk_text, @@ -2124,7 +2174,8 @@ impl StreamingProcessor { } // Tool call handling: incremental streaming parser - if !in_reasoning && streaming_tool_parser.is_some() { + let tool_parser_active = !in_reasoning && streaming_tool_parser.is_some(); + if tool_parser_active { if is_specific_function { // Specific function: entire output is arguments for one tool if !has_tool_calls { @@ -2176,7 +2227,7 @@ impl StreamingProcessor { &MessageStreamEvent::ContentBlockDelta { index: current_block_index, delta: ContentBlockDelta::InputJsonDelta { - partial_json: normal_text, + partial_json: std::mem::take(&mut normal_text), }, }, ) @@ -2291,11 +2342,10 @@ impl StreamingProcessor { } } } - continue; } // Regular text emission (no tools active) - if !normal_text.is_empty() { + if !tool_parser_active && !normal_text.is_empty() { if !text_block_open { Self::send_messages_event( tx, @@ -2321,6 +2371,12 @@ impl StreamingProcessor { ) .await?; } + + if router_string_stop { + prompt_tokens = chunk.prompt_tokens(); + router_terminated = true; + break; + } } ProtoResponseVariant::Complete(complete) => { // Flush stop decoder @@ -2513,8 +2569,9 @@ impl StreamingProcessor { // Phase 5: Emit message_stop Self::send_messages_event(tx, &mut sse_buffer, &MessageStreamEvent::MessageStop).await?; - // Mark stream completed - grpc_stream.mark_completed(); + if !router_terminated { + grpc_stream.mark_completed(); + } if let Some(handle) = reservation { if saw_complete { @@ -2844,6 +2901,10 @@ impl StreamingProcessor { let mut stop_decoders: HashMap = HashMap::new(); let mut is_firsts: HashMap = HashMap::new(); let mut stopped_indices: HashSet = HashSet::new(); + let mut terminal_indices: HashSet = HashSet::new(); + let expected_choices = completion_request.n.unwrap_or(1).max(1) as usize; + let mut has_router_stop = false; + let mut router_terminated = false; let mut sse_buffer = Vec::with_capacity(512); let mut chunk_text = String::new(); // For n>1, each index shares the same prompt — use max across Complete @@ -2897,6 +2958,9 @@ impl StreamingProcessor { let (decoded_text, stopped) = Self::process_chunk_tokens(stop_decoder, chunk.token_ids()); + // `matched_stop` is only set for router-matched string stops; + // token-level stops leave the backend to terminate itself. + let matched_sequence = stopped && stop_decoder.matched_stop().is_some(); chunk_text.clear(); chunk_text.push_str(&decoded_text); @@ -2934,6 +2998,11 @@ impl StreamingProcessor { // the backend's eventual Complete carries "length". This is // intentional — the local stop sequence fired first. stopped_indices.insert(index); + terminal_indices.insert(index); + has_router_stop |= matched_sequence; + total_prompt = total_prompt.max(chunk.prompt_tokens()); + total_cached = total_cached.max(chunk.cached_tokens()); + reasoning_tokens.insert(index, chunk.reasoning_tokens()); if let Some(sfx) = suffix { let suffix_chunk = CompletionStreamResponse { @@ -2974,6 +3043,11 @@ impl StreamingProcessor { tx.send(Ok(Bytes::from(sse_buffer.clone()))) .await .map_err(|_| "Channel closed".to_string())?; + + if has_router_stop && terminal_indices.len() >= expected_choices { + router_terminated = true; + break; + } } } ProtoResponseVariant::Complete(complete) => { @@ -3104,12 +3178,20 @@ impl StreamingProcessor { tx.send(Ok(Bytes::from(sse_buffer.clone()))) .await .map_err(|_| "Channel closed".to_string())?; + + terminal_indices.insert(index); + if has_router_stop && terminal_indices.len() >= expected_choices { + router_terminated = true; + break; + } } ProtoResponseVariant::None => continue, } } - grpc_stream.mark_completed(); + if !router_terminated { + grpc_stream.mark_completed(); + } // `completed_indices` counts distinct indices that finished cleanly // via a `Complete` message. A clean EOF partway through this unit's