Skip to main content

bge_m3_embedding_server/embedder/worker/
run.rs

1// Copyright (c) 2026 J. Patrick Fulton
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Blocking worker thread and request dispatch loop.
16
17use std::path::PathBuf;
18use std::sync::Arc;
19use std::sync::atomic::{AtomicUsize, Ordering};
20
21use anyhow::Result;
22use tokio::runtime::Handle;
23use tokio::sync::{Mutex, mpsc};
24use tracing::info;
25
26use super::config::WorkerConfig;
27use super::dispatch::{DispatchOutcome, dispatch_request, reply_request_load_error};
28use super::guard::{
29    InferenceOutcome, WorkerGuard, build_shape_guard, next_consecutive_failures,
30    should_unload_on_outcome,
31};
32use super::probe::probe_run_dense;
33use super::propagation::drain_engine_propagation;
34use super::startup::{StartupOutcome, startup_worker};
35use crate::config::EpSelection;
36use crate::embedder::session::{GpuSessionConfig, load_models};
37use crate::embedder::trt_warmup::trt_prewarm;
38use crate::embedder::types::EmbedRequest;
39
40#[allow(clippy::needless_pass_by_value, clippy::too_many_lines)]
41pub(in crate::embedder) fn run_worker(
42    id: usize,
43    cache_dir: PathBuf,
44    rx: Arc<Mutex<mpsc::Receiver<EmbedRequest>>>,
45    ready_tx: mpsc::Sender<Result<usize>>,
46    live_workers: Arc<AtomicUsize>,
47    loaded_workers: Arc<AtomicUsize>,
48    config: WorkerConfig,
49) -> Result<()> {
50    let _guard = WorkerGuard(Arc::clone(&live_workers));
51
52    let rt = Handle::current();
53    let StartupOutcome {
54        initial_models,
55        detected_sm,
56    } = startup_worker(id, &cache_dir, &ready_tx, &config, &rt)?;
57    let mut models: Option<(ort::session::Session, tokenizers::Tokenizer)> = Some(initial_models);
58
59    // Tracks shapes already prewarmed by this worker via engine propagation so
60    // we skip shapes we originated (which are already in our TRT profile after
61    // the originating worker inserts into warmed_local before broadcasting).
62    let mut warmed_local: std::collections::HashSet<(usize, usize)> =
63        std::collections::HashSet::new();
64
65    // Derive per-worker broadcast receiver from the shared sender in config.
66    // Each call to tx.subscribe() creates an independent receiver starting from
67    // the current channel position, so workers only see shapes broadcast after
68    // their own subscribe point (i.e. after model load is complete).
69    let mut engine_propagation_rx = config
70        .engine_propagation_tx
71        .as_ref()
72        .map(tokio::sync::broadcast::Sender::subscribe);
73
74    // Per-worker consecutive-failure counter for the inference circuit breaker.
75    // Incremented on every inference Err; reset to 0 on every inference Ok.
76    // When it reaches `config.circuit_breaker_threshold` the worker unloads
77    // its models (drops the ORT session, clears the CUDA arena) and decrements
78    // `loaded_workers`. The standard idle-reload path handles model recovery on
79    // the next incoming request.
80    let mut consecutive_failures: u64 = 0;
81
82    info!("Worker {id} entering request loop");
83    loop {
84        // Drain peer engine-ready notifications and run trt_prewarm for any
85        // new shapes so the in-memory TRT profile is extended before the next
86        // real request for that shape arrives (~1-3s fast disk-load vs. full JIT).
87        if config.ep == EpSelection::TensorRt
88            && let Some(ref mut bcast_rx) = engine_propagation_rx
89            && let Some((session, _)) = models.as_mut()
90        {
91            let sm = detected_sm.as_deref();
92            let ceiling = &config.warmed_seq_ceiling;
93            drain_engine_propagation(bcast_rx, &mut warmed_local, id, |shape| {
94                let started = std::time::Instant::now();
95                let stats = trt_prewarm(session, &[shape], id, &cache_dir, sm);
96                let elapsed_ms = u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX);
97                if stats.warmed > 0 || stats.fully_cached {
98                    // A peer-propagated plan now covers this shape on disk and
99                    // in this worker's session: extend the guard ceiling.
100                    ceiling.fetch_max(stats.max_warmed_seq.max(shape.1), Ordering::AcqRel);
101                }
102                tracing::info!(
103                    target: "bge_m3_embedding_server::trt_shape",
104                    worker_id = id,
105                    chunk_batch = shape.0,
106                    chunk_max_seq = shape.1,
107                    elapsed_ms,
108                    warmed = stats.warmed,
109                    fully_cached = stats.fully_cached,
110                    detected_sm = sm.unwrap_or("unfiltered"),
111                    "engine_propagation_complete"
112                );
113            });
114        }
115
116        let msg = if let Some(timeout) = config.idle_timeout.filter(|_| models.is_some()) {
117            rt.block_on(async {
118                tokio::time::timeout(timeout, async { rx.lock().await.recv().await }).await
119            })
120        } else {
121            rt.block_on(async { Ok(rx.lock().await.recv().await) })
122        };
123
124        match msg {
125            Err(_elapsed) => {
126                models = None;
127                loaded_workers.fetch_sub(1, Ordering::AcqRel);
128                tracing::info!("Worker {id} unloaded models after idle timeout");
129            }
130            Ok(None) => {
131                if models.is_some() {
132                    loaded_workers.fetch_sub(1, Ordering::AcqRel);
133                }
134                info!("Worker {id} channel closed, shutting down");
135                break;
136            }
137            Ok(Some(request)) => {
138                if models.is_none() {
139                    tracing::info!("Worker {id} reloading models after idle...");
140                    let reload_start = std::time::Instant::now();
141                    match load_models(
142                        &GpuSessionConfig {
143                            cache_dir: &cache_dir,
144                            model_variant: config.model_variant,
145                            max_seq_length: config.max_seq_length,
146                            intra_threads: config.intra_threads,
147                            ep: config.ep,
148                            device_id: config.device_id,
149                            trt_max_workspace_bytes: config.trt_max_workspace_bytes,
150                            gpu_mem_limit_bytes: config.gpu_mem_limit_bytes,
151                        },
152                        false,
153                    ) {
154                        Ok(mut m) => {
155                            // Prime the freshly-loaded session arena so the
156                            // first incoming request after idle reload doesn't
157                            // pay the ~1 GiB lazy-arena-init cost. Same
158                            // rationale as the startup priming in the
159                            // load-models Ok arm above.
160                            let prime_ids = ndarray::Array2::<i64>::zeros((1, 8));
161                            let prime_mask = ndarray::Array2::<i64>::ones((1, 8));
162                            if let Err(e) = probe_run_dense(&mut m.0, &prime_ids, &prime_mask) {
163                                tracing::warn!(
164                                    error = %e,
165                                    "Worker {id} post-reload arena prime failed"
166                                );
167                            }
168                            models = Some(m);
169                            loaded_workers.fetch_add(1, Ordering::AcqRel);
170                            tracing::info!(
171                                elapsed_ms = reload_start.elapsed().as_millis(),
172                                "Worker {id} reloaded models"
173                            );
174                        }
175                        Err(e) => {
176                            tracing::error!(error = %e, "Worker {id} failed to reload models");
177                            let err = anyhow::anyhow!("Model reload failed: {e}");
178                            reply_request_load_error(request, err);
179                            continue;
180                        }
181                    }
182                }
183
184                // In-band TRT JIT admission guard, rebuilt per request from the
185                // live pool-wide warmed-seq ceiling. `None` on non-TRT EPs or
186                // when the guard is disabled. Passed into the embed functions
187                // which refuse dangerous, uncovered chunk shapes before
188                // `session.run()` (returning a TrtJitRejection → HTTP 503).
189                let shape_guard = build_shape_guard(&config);
190
191                // Flags for circuit-breaker and fatal-exit decisions; set
192                // inside the borrow scope of `session`/`tokenizer` and acted
193                // on AFTER that scope ends so `models` can be safely mutated.
194                // `InferenceOutcome::TrtFatal` triggers a worker exit.
195                // `InferenceOutcome::CircuitBreak` unloads the ORT session.
196                let outcome: InferenceOutcome;
197                let skip_to_next: bool;
198
199                {
200                    let (session, tokenizer) =
201                        models.as_mut().expect("models loaded after reload check");
202
203                    let DispatchOutcome {
204                        outcome: inner_outcome,
205                        skip: inner_skip,
206                    } = dispatch_request(
207                        request,
208                        session,
209                        tokenizer,
210                        &config,
211                        id,
212                        &cache_dir,
213                        detected_sm.as_deref(),
214                        &mut warmed_local,
215                        consecutive_failures,
216                        shape_guard.as_ref(),
217                    );
218                    outcome = inner_outcome;
219                    skip_to_next = inner_skip;
220                } // end borrow scope — session and tokenizer dropped here
221
222                // --- Post-inference actions (models reborrow is now safe) ---
223                if skip_to_next {
224                    continue;
225                }
226                consecutive_failures = next_consecutive_failures(outcome, consecutive_failures);
227                if matches!(outcome, InferenceOutcome::TrtFatal) {
228                    return Err(anyhow::anyhow!(
229                        "Worker {id} hit fatal TRT engine build error; exiting"
230                    ));
231                }
232                if should_unload_on_outcome(outcome) {
233                    models = None;
234                    loaded_workers.fetch_sub(1, Ordering::AcqRel);
235                }
236            }
237        }
238    }
239
240    Ok(())
241}