diff --git a/sgl-model-gateway/src/core/worker.rs b/sgl-model-gateway/src/core/worker.rs index 006b12485..53dbe0186 100644 --- a/sgl-model-gateway/src/core/worker.rs +++ b/sgl-model-gateway/src/core/worker.rs @@ -561,6 +561,10 @@ impl BasicWorker { Ok(self.url()) } } + + fn update_running_requests_metrics(&self) { + RouterMetrics::set_running_requests(self.url(), self.load()); + } } #[async_trait] @@ -631,6 +635,7 @@ impl Worker for BasicWorker { fn increment_load(&self) { self.load_counter.fetch_add(1, Ordering::Relaxed); + self.update_running_requests_metrics(); } fn decrement_load(&self) { @@ -639,10 +644,12 @@ impl Worker for BasicWorker { current.checked_sub(1) }) .ok(); + self.update_running_requests_metrics(); } fn reset_load(&self) { self.load_counter.store(0, Ordering::Relaxed); + self.update_running_requests_metrics(); } fn processed_requests(&self) -> usize { diff --git a/sgl-model-gateway/src/routers/http/router.rs b/sgl-model-gateway/src/routers/http/router.rs index 9ae55ce97..162297d91 100644 --- a/sgl-model-gateway/src/routers/http/router.rs +++ b/sgl-model-gateway/src/routers/http/router.rs @@ -234,7 +234,7 @@ impl Router { }; let load_incremented = if policy.name() == "cache_aware" { - increment_load(&worker); + worker.increment_load(); true } else { false @@ -274,7 +274,7 @@ impl Router { // won't have done it (it only decrements on success or non-retryable failures) if is_retryable_status(response.status()) && load_incremented { if let Some(cleanup_worker) = worker_for_cleanup { - decrement_load(&cleanup_worker); + cleanup_worker.decrement_load(); } } @@ -507,7 +507,7 @@ impl Router { // Decrement load on error if it was incremented if load_incremented { if let Some(ref w) = worker { - decrement_load(w); + w.decrement_load(); } } @@ -533,7 +533,7 @@ impl Router { // IMPORTANT: Decrement load on error before returning if load_incremented { if let Some(ref w) = worker { - decrement_load(w); + w.decrement_load(); } } @@ -545,7 +545,7 @@ impl Router { // Decrement load counter for non-streaming requests if it was incremented if load_incremented { if let Some(ref w) = worker { - decrement_load(w); + w.decrement_load(); } } @@ -573,7 +573,7 @@ impl Router { // Check for stream end marker using memmem for efficiency if memmem::find(&bytes, b"data: [DONE]").is_some() { if let Some(ref w) = stream_worker { - decrement_load(w); + w.decrement_load(); decremented = true; } } @@ -589,7 +589,7 @@ impl Router { } if !decremented { if let Some(ref w) = stream_worker { - decrement_load(w); + w.decrement_load(); } } }); @@ -659,16 +659,6 @@ impl Router { } } -fn increment_load(w: &Arc) { - w.increment_load(); - RouterMetrics::set_running_requests(w.url(), w.load()); -} - -fn decrement_load(w: &Arc) { - w.decrement_load(); - RouterMetrics::set_running_requests(w.url(), w.load()); -} - fn convert_reqwest_error(e: reqwest::Error) -> Response { let url = e .url()