@@ -24,7 +24,7 @@ use stargate_protocol::common::{
2424use stargate_protocol:: tunnel_contract:: HEADER_STARGATE_EXPECTED_QUEUE_MS ;
2525
2626use crate :: request_observer:: { RequestObservation , RequestObservationState , RequiredTunnelHeaders } ;
27- use crate :: runtime_state:: ModelGeneration ;
27+ use crate :: runtime_state:: { ModelGeneration , RequestInstance } ;
2828
2929pub ( crate ) const RETRY_REASON_QUEUE_ESTIMATE_MISMATCH : & str = "queue_estimate_mismatch" ;
3030
@@ -55,10 +55,16 @@ pub(crate) struct LiveRequestState {
5555
5656#[ derive( Debug , Default ) ]
5757struct QueueAdmissionState {
58- requests : HashMap < String , LiveRequest > ,
58+ requests : HashMap < String , LiveRequestRecord > ,
5959 models : HashMap < ModelGeneration , QueueModelState > ,
6060}
6161
62+ #[ derive( Debug ) ]
63+ struct LiveRequestRecord {
64+ instance : Option < RequestInstance > ,
65+ request : LiveRequest ,
66+ }
67+
6268#[ derive( Debug , Default ) ]
6369struct QueueModelState {
6470 last_mean_input_tps : Option < f64 > ,
@@ -155,6 +161,7 @@ pub(crate) struct QueueModelSnapshot {
155161pub ( crate ) struct QueueTrackedRequestGuard {
156162 live_requests : LiveRequestState ,
157163 request_id : String ,
164+ instance : RequestInstance ,
158165 finished : bool ,
159166}
160167
@@ -221,7 +228,7 @@ impl LiveRequestState {
221228 . lock ( )
222229 . requests
223230 . get ( request_id)
224- . map ( |request| request. generation ( ) . clone ( ) )
231+ . map ( |record| record . request . generation ( ) . clone ( ) )
225232 }
226233
227234 pub ( crate ) fn update_generation_throughput (
@@ -247,7 +254,7 @@ impl LiveRequestState {
247254 let request_ids = state
248255 . requests
249256 . iter ( )
250- . filter ( |( _, request ) | request. generation ( ) == generation)
257+ . filter ( |( _, record ) | record . request . generation ( ) == generation)
251258 . map ( |( request_id, _) | request_id. clone ( ) )
252259 . collect :: < Vec < _ > > ( ) ;
253260 for request_id in request_ids {
@@ -302,7 +309,7 @@ impl LiveRequestState {
302309 let excluded_request = state
303310 . requests
304311 . get ( & required. request_id )
305- . and_then ( |request | match request {
312+ . and_then ( |record | match & record . request {
306313 LiveRequest :: Queue ( queue, _) => Some ( queue) ,
307314 LiveRequest :: Observed ( _) => None ,
308315 } )
@@ -343,11 +350,11 @@ impl LiveRequestState {
343350 self . track_generation_request ( required, ModelGeneration :: new ( required. model_id . clone ( ) , 0 ) )
344351 }
345352
346- pub ( crate ) fn track_generation_request (
353+ pub ( crate ) fn begin_request (
347354 & self ,
348355 required : & RequiredTunnelHeaders ,
349356 generation : ModelGeneration ,
350- ) -> QueueTrackedRequestGuard {
357+ ) {
351358 let request_id = required. request_id . clone ( ) ;
352359 let request = TrackedPromptRequest {
353360 generation,
@@ -360,12 +367,40 @@ impl LiveRequestState {
360367 let mut state = self . inner . lock ( ) ;
361368 let observed = state
362369 . remove_request ( & request_id)
363- . and_then ( |( _, request) | request. into_observed ( ) ) ;
364- state. insert_request ( request_id. clone ( ) , LiveRequest :: Queue ( request, observed) ) ;
370+ . and_then ( |( _, request) | request. request . into_observed ( ) ) ;
371+ state. insert_request (
372+ request_id,
373+ LiveRequest :: Queue ( request, observed) ,
374+ Some ( required. request_instance . clone ( ) ) ,
375+ ) ;
365376 }
377+ }
378+
379+ #[ cfg( test) ]
380+ pub ( crate ) fn track_generation_request (
381+ & self ,
382+ required : & RequiredTunnelHeaders ,
383+ generation : ModelGeneration ,
384+ ) -> QueueTrackedRequestGuard {
385+ self . begin_request ( required, generation) ;
386+ self . request_guard ( required)
387+ }
388+
389+ pub ( crate ) fn track_existing_request (
390+ & self ,
391+ required : & RequiredTunnelHeaders ,
392+ ) -> Option < QueueTrackedRequestGuard > {
393+ self . inner
394+ . lock ( )
395+ . owns_instance ( & required. request_id , & required. request_instance )
396+ . then ( || self . request_guard ( required) )
397+ }
398+
399+ fn request_guard ( & self , required : & RequiredTunnelHeaders ) -> QueueTrackedRequestGuard {
366400 QueueTrackedRequestGuard {
367401 live_requests : self . clone ( ) ,
368- request_id,
402+ request_id : required. request_id . clone ( ) ,
403+ instance : required. request_instance . clone ( ) ,
369404 finished : false ,
370405 }
371406 }
@@ -390,22 +425,30 @@ impl LiveRequestState {
390425 observe : impl FnOnce ( & RequestObservationTransition ) ,
391426 ) -> RequestObservationTransition {
392427 let generation = ModelGeneration :: new ( observation. model_id . clone ( ) , 0 ) ;
393- self . transition_generation_observation_with ( observation, Some ( & generation) , observe)
428+ self . transition_generation_observation_with ( observation, Some ( & generation) , None , observe)
429+ . unwrap ( )
394430 }
395431
396432 pub ( crate ) fn transition_generation_observation_with (
397433 & self ,
398434 observation : & RequestObservation ,
399435 generation : Option < & ModelGeneration > ,
436+ instance : Option < & RequestInstance > ,
400437 observe : impl FnOnce ( & RequestObservationTransition ) ,
401- ) -> RequestObservationTransition {
438+ ) -> Option < RequestObservationTransition > {
402439 let _order = self . observation_order . lock ( ) ;
403- let transition = self
404- . inner
405- . lock ( )
406- . transition_observation ( observation, generation) ;
440+ let transition = {
441+ let mut state = self . inner . lock ( ) ;
442+ if let Some ( instance) = instance
443+ && generation. is_some ( )
444+ && !state. owns_instance ( & observation. request_id , instance)
445+ {
446+ return None ;
447+ }
448+ state. transition_observation ( observation, generation, instance)
449+ } ;
407450 observe ( & transition) ;
408- transition
451+ Some ( transition)
409452 }
410453
411454 pub ( crate ) fn update_active_output_tps (
@@ -418,24 +461,44 @@ impl LiveRequestState {
418461 . update_active_output_tps ( request_id, active_chat_output_tps)
419462 }
420463
421- pub ( crate ) fn finish_queue_request ( & self , request_id : & str ) {
464+ pub ( crate ) fn finish_queue_request (
465+ & self ,
466+ request_id : & str ,
467+ instance : Option < & RequestInstance > ,
468+ ) {
422469 let mut state = self . inner . lock ( ) ;
423- if let Some ( ( request_id, request) ) = state. remove_request ( request_id)
424- && let Some ( observed) = request. into_observed ( )
470+ if instance. is_some_and ( |instance| !state. owns_instance ( request_id, instance) ) {
471+ return ;
472+ }
473+ if let Some ( ( request_id, record) ) = state. remove_request ( request_id)
474+ && let Some ( observed) = record. request . into_observed ( )
425475 {
426- state. insert_request ( request_id, LiveRequest :: Observed ( observed) ) ;
476+ state. insert_request ( request_id, LiveRequest :: Observed ( observed) , record . instance ) ;
427477 }
428478 }
429479}
430480
431481impl QueueAdmissionState {
482+ fn owns_instance ( & self , request_id : & str , instance : & RequestInstance ) -> bool {
483+ self . requests
484+ . get ( request_id)
485+ . is_some_and ( |record| record. instance . as_ref ( ) == Some ( instance) )
486+ }
487+
432488 fn transition_observation (
433489 & mut self ,
434490 observation : & RequestObservation ,
435491 generation : Option < & ModelGeneration > ,
492+ instance : Option < & RequestInstance > ,
436493 ) -> RequestObservationTransition {
494+ let prior = self . remove_request ( & observation. request_id ) ;
495+ let instance = instance. cloned ( ) . or_else ( || {
496+ prior
497+ . as_ref ( )
498+ . and_then ( |( _, record) | record. instance . clone ( ) )
499+ } ) ;
437500 let ( request_id, prior_queue, prior_observed) =
438- match self . remove_request ( & observation . request_id ) {
501+ match prior . map ( | ( id , record ) | ( id , record . request ) ) {
439502 Some ( ( request_id, LiveRequest :: Queue ( queue, observed) ) ) => {
440503 ( request_id, Some ( queue) , observed)
441504 }
@@ -493,6 +556,7 @@ impl QueueAdmissionState {
493556 } ,
494557 Some ( current. clone ( ) ) ,
495558 ) ,
559+ instance,
496560 ) ;
497561 Some ( current)
498562 }
@@ -513,14 +577,28 @@ impl QueueAdmissionState {
513577 request_id : & str ,
514578 active_chat_output_tps : Option < f64 > ,
515579 ) -> Option < String > {
516- let ( request_id, mut request) = self . remove_request ( request_id) ?;
517- let model_id = request. update_active_output_tps ( active_chat_output_tps) ;
518- self . insert_request ( request_id, request) ;
580+ let ( request_id, mut record) = self . remove_request ( request_id) ?;
581+ let model_id = record
582+ . request
583+ . update_active_output_tps ( active_chat_output_tps) ;
584+ self . insert_request ( request_id, record. request , record. instance ) ;
519585 model_id
520586 }
521587
522- fn advance_request_phase ( & mut self , request_id : & str , next_phase : TrackedPromptPhase ) {
523- let Some ( LiveRequest :: Queue ( request, _) ) = self . requests . get_mut ( request_id) else {
588+ fn advance_request_phase (
589+ & mut self ,
590+ request_id : & str ,
591+ instance : & RequestInstance ,
592+ next_phase : TrackedPromptPhase ,
593+ ) {
594+ if !self . owns_instance ( request_id, instance) {
595+ return ;
596+ }
597+ let Some ( LiveRequest :: Queue ( request, _) ) = self
598+ . requests
599+ . get_mut ( request_id)
600+ . map ( |record| & mut record. request )
601+ else {
524602 return ;
525603 } ;
526604 if next_phase <= request. phase {
@@ -535,15 +613,21 @@ impl QueueAdmissionState {
535613 request. phase = next_phase;
536614 }
537615
538- fn remove_request ( & mut self , request_id : & str ) -> Option < ( String , LiveRequest ) > {
539- let ( request_id, request ) = self . requests . remove_entry ( request_id) ?;
540- self . adjust_live_request ( & request, -1 ) ;
541- Some ( ( request_id, request ) )
616+ fn remove_request ( & mut self , request_id : & str ) -> Option < ( String , LiveRequestRecord ) > {
617+ let ( request_id, record ) = self . requests . remove_entry ( request_id) ?;
618+ self . adjust_live_request ( & record . request , -1 ) ;
619+ Some ( ( request_id, record ) )
542620 }
543621
544- fn insert_request ( & mut self , request_id : String , request : LiveRequest ) {
622+ fn insert_request (
623+ & mut self ,
624+ request_id : String ,
625+ request : LiveRequest ,
626+ instance : Option < RequestInstance > ,
627+ ) {
545628 self . adjust_live_request ( & request, 1 ) ;
546- self . requests . insert ( request_id, request) ;
629+ self . requests
630+ . insert ( request_id, LiveRequestRecord { instance, request } ) ;
547631 }
548632
549633 fn adjust_live_request ( & mut self , request : & LiveRequest , delta : i8 ) {
@@ -786,17 +870,26 @@ impl TrackedPromptPhase {
786870impl QueueTrackedRequestGuard {
787871 pub ( crate ) fn on_backend_submission ( & mut self ) {
788872 let mut state = self . live_requests . inner . lock ( ) ;
789- state. advance_request_phase ( & self . request_id , TrackedPromptPhase :: InputProcessing ) ;
873+ state. advance_request_phase (
874+ & self . request_id ,
875+ & self . instance ,
876+ TrackedPromptPhase :: InputProcessing ,
877+ ) ;
790878 }
791879
792880 pub ( crate ) fn observe_output ( & mut self ) {
793881 let mut state = self . live_requests . inner . lock ( ) ;
794- state. advance_request_phase ( & self . request_id , TrackedPromptPhase :: OutputGeneration ) ;
882+ state. advance_request_phase (
883+ & self . request_id ,
884+ & self . instance ,
885+ TrackedPromptPhase :: OutputGeneration ,
886+ ) ;
795887 }
796888
797889 pub ( crate ) fn finish ( & mut self ) {
798890 if !self . finished {
799- self . live_requests . finish_queue_request ( & self . request_id ) ;
891+ self . live_requests
892+ . finish_queue_request ( & self . request_id , Some ( & self . instance ) ) ;
800893 self . finished = true ;
801894 }
802895 }
@@ -840,7 +933,7 @@ mod tests {
840933 . lock ( )
841934 . requests
842935 . values ( )
843- . filter ( |request | matches ! ( request, LiveRequest :: Queue ( ..) ) )
936+ . filter ( |record | matches ! ( record . request, LiveRequest :: Queue ( ..) ) )
844937 . count ( )
845938 }
846939 }
@@ -856,6 +949,7 @@ mod tests {
856949 input_tokens : u64 ,
857950 ) -> RequiredTunnelHeaders {
858951 RequiredTunnelHeaders {
952+ request_instance : Default :: default ( ) ,
859953 request_id : request_id. to_string ( ) ,
860954 routing_key : None ,
861955 model_id : model_id. to_string ( ) ,
@@ -910,6 +1004,27 @@ mod tests {
9101004 }
9111005 }
9121006
1007+ #[ test]
1008+ fn stale_guards_cannot_advance_or_remove_reused_request_ids ( ) {
1009+ for replacement_model in [ "model-a" , "model-b" ] {
1010+ let live = LiveRequestState :: default ( ) ;
1011+ let mut first = live. track_request ( & required_for_model ( "same-id" , "model-a" , 0 , 100 ) ) ;
1012+ let replacement =
1013+ live. track_request ( & required_for_model ( "same-id" , replacement_model, 1 , 20 ) ) ;
1014+ let expected = live. snapshot_model ( replacement_model) ;
1015+ first. on_backend_submission ( ) ;
1016+ first. observe_output ( ) ;
1017+ drop ( first) ;
1018+ assert_eq ! ( live. snapshot_model( replacement_model) , expected) ;
1019+ assert_eq ! ( expected. num_running_queries, 1 ) ;
1020+ drop ( replacement) ;
1021+ assert_eq ! (
1022+ live. snapshot_model( replacement_model) . num_running_queries,
1023+ 0
1024+ ) ;
1025+ }
1026+ }
1027+
9131028 #[ test]
9141029 fn one_live_request_transition_updates_queue_and_active_output_load ( ) {
9151030 let live_requests = LiveRequestState :: default ( ) ;
0 commit comments