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}