1use super::*;
2
3#[derive(Debug, Clone)]
4pub struct SessionManagerUpdate {
5 pub session_id: String,
6 pub view: ManagedSessionView,
7}
8
9pub(crate) type SessionCpuTable =
10 mj_core::snapshot_map::SnapshotMap<String, mj_client::runtime_feed::SessionCpuView>;
11
12pub struct SessionManagerChannels {
13 pub session_cpu: watch::Receiver<SessionCpuTable>,
14 pub targets: watch::Sender<Vec<RelaySessionTarget>>,
15 pub control: SessionManagerControl,
16 pub updates: SessionManagerUpdates,
17 pub shutdown: SessionManagerShutdown,
18}
19
20pub struct RemoteSessionManagerChannels {
26 pub targets: watch::Sender<Vec<RelaySessionTarget>>,
27 pub control: SessionManagerControl,
28 pub updates: SessionManagerUpdates,
29 pub shutdown: SessionManagerShutdown,
30 pub publisher: RemoteSessionPublisher,
31 pub requests: RemoteSessionRequests,
32}
33
34#[derive(Clone)]
35pub struct RemoteSessionPublisher {
36 pub(super) updates: mpsc::UnboundedSender<RemoteManagerUpdate>,
37}
38
39impl RemoteSessionPublisher {
40 pub async fn publish(&self, session_id: String, view: ManagedSessionView) -> Result<()> {
41 self.updates
42 .send(RemoteManagerUpdate::Publish { session_id, view })
43 .context("remote session manager stopped")
44 }
45
46 pub fn try_publish(&self, session_id: String, view: ManagedSessionView) -> Result<()> {
47 self.updates
48 .send(RemoteManagerUpdate::Publish { session_id, view })
49 .context("remote session manager update queue is unavailable")
50 }
51}
52
53pub struct RemoteSessionRequests {
54 pub(super) requests: mpsc::Receiver<RemoteSessionRequest>,
55}
56
57impl RemoteSessionRequests {
58 pub async fn recv(&mut self) -> Option<RemoteSessionRequest> {
59 self.requests.recv().await
60 }
61}
62
63pub enum RemoteSessionRequest {
64 Submit {
65 session_id: String,
66 command_id: String,
67 command: RelayCommand,
68 admission: Option<ReviewDeliveryAdmission>,
69 reply: oneshot::Sender<std::result::Result<u64, mj_client::session::SubmitFailure>>,
70 },
71 Sync {
72 session_id: String,
73 reply: oneshot::Sender<std::result::Result<(), String>>,
74 },
75 RespondElicitation {
76 session_id: String,
77 elicitation_id: String,
78 response: ElicitationResponse,
79 reply: oneshot::Sender<std::result::Result<(), String>>,
80 },
81 StopBackgroundTask {
82 session_id: String,
83 background_task_id: String,
84 reply: oneshot::Sender<std::result::Result<(), String>>,
85 },
86 Reviewer {
87 session_id: String,
88 role: Option<String>,
90 action: ReviewerAction,
91 reply: oneshot::Sender<std::result::Result<ReviewerOutcome, String>>,
92 },
93}
94
95impl RemoteSessionRequest {
96 pub fn session_id(&self) -> &str {
99 match self {
100 Self::Submit { session_id, .. }
101 | Self::Sync { session_id, .. }
102 | Self::RespondElicitation { session_id, .. }
103 | Self::StopBackgroundTask { session_id, .. }
104 | Self::Reviewer { session_id, .. } => session_id,
105 }
106 }
107}
108
109const REQUESTS_PER_STREAM: usize = 32;
111const REQUESTS_TOTAL: usize = 256;
112type ForwardRequest = std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send>>;
113type RequestBudgets = Arc<Mutex<std::collections::HashMap<SessionRequestStream, usize>>>;
114
115pub struct SessionRequestOrder {
121 sender: mpsc::Sender<OrderedRequest>,
122 budgets: RequestBudgets,
123 supervisor: tokio::task::JoinHandle<()>,
124}
125
126#[derive(Debug, Clone, PartialEq, Eq, Hash)]
127pub(super) enum SessionRequestStream {
128 Primary(String),
129 Reviewer(String, Option<String>),
130}
131
132struct RequestPermit {
133 budgets: RequestBudgets,
134 stream: SessionRequestStream,
135}
136
137impl Drop for RequestPermit {
138 fn drop(&mut self) {
139 let mut budgets = self.budgets.lock().expect("request budgets poisoned");
140 let count = budgets.get_mut(&self.stream).expect("admitted request");
141 *count -= 1;
142 if *count == 0 {
143 budgets.remove(&self.stream);
144 }
145 }
146}
147
148struct OrderedRequest {
149 permit: RequestPermit,
150 forward: ForwardRequest,
151}
152
153impl RemoteSessionRequest {
154 pub(crate) fn reject(self, message: &str) {
155 match self {
156 Self::Submit { reply, .. } => {
157 let _ = reply.send(Err(message.to_owned().into()));
158 }
159 Self::Sync { reply, .. }
160 | Self::RespondElicitation { reply, .. }
161 | Self::StopBackgroundTask { reply, .. } => {
162 let _ = reply.send(Err(message.to_owned()));
163 }
164 Self::Reviewer { reply, .. } => {
165 let _ = reply.send(Err(message.to_owned()));
166 }
167 }
168 }
169}
170
171impl Default for SessionRequestOrder {
172 fn default() -> Self {
173 Self::new()
174 }
175}
176
177impl SessionRequestOrder {
178 #[must_use]
179 pub fn new() -> Self {
180 let (sender, receiver) = mpsc::channel(REQUESTS_TOTAL);
181 Self {
182 sender,
183 budgets: Default::default(),
184 supervisor: tokio::spawn(supervise_requests(receiver)),
185 }
186 }
187
188 pub fn dispatch<F, Fut>(&mut self, request: RemoteSessionRequest, forward: F)
190 where
191 F: FnOnce(RemoteSessionRequest) -> Fut + Send + 'static,
192 Fut: std::future::Future<Output = ()> + Send + 'static,
193 {
194 let stream = match &request {
195 RemoteSessionRequest::Reviewer {
196 session_id, role, ..
197 } => SessionRequestStream::Reviewer(session_id.clone(), role.clone()),
198 _ => SessionRequestStream::Primary(request.session_id().to_owned()),
199 };
200 let mut budgets = self.budgets.lock().expect("request budgets poisoned");
201 if budgets.get(&stream).copied().unwrap_or(0) >= REQUESTS_PER_STREAM
202 || budgets.values().sum::<usize>() >= REQUESTS_TOTAL
203 {
204 drop(budgets);
205 request.reject("session request queue is full; request was not accepted");
206 return;
207 }
208 let Ok(slot) = self.sender.try_reserve() else {
211 drop(budgets);
212 request.reject("session request dispatcher is unavailable; request was not accepted");
213 return;
214 };
215 *budgets.entry(stream.clone()).or_default() += 1;
216 drop(budgets);
217 slot.send(OrderedRequest {
218 permit: RequestPermit {
219 budgets: self.budgets.clone(),
220 stream,
221 },
222 forward: Box::pin(async move { forward(request).await }),
223 });
224 }
225
226 pub async fn drain(self) -> Result<()> {
228 drop(self.sender);
229 self.supervisor
230 .await
231 .context("session request supervisor failed")
232 }
233}
234
235async fn supervise_requests(mut requests: mpsc::Receiver<OrderedRequest>) {
236 let mut tasks = tokio::task::JoinSet::new();
237 let mut running = std::collections::HashMap::new();
238 let mut pending =
239 std::collections::HashMap::<SessionRequestStream, VecDeque<OrderedRequest>>::new();
240 let mut closed = false;
241 loop {
242 tokio::select! {
243 biased;
246 completed = tasks.join_next_with_id(), if !tasks.is_empty() => {
247 let task_id = match completed.expect("nonempty request tasks") {
248 Ok((id, ())) => id,
249 Err(error) => {
250 tracing::error!(stream = ?running.get(&error.id()), %error, "session request task failed");
251 error.id()
252 }
253 };
254 let stream = running.remove(&task_id).expect("registered request task");
255 if let Some(queue) = pending.get_mut(&stream) {
256 if let Some(request) = queue.pop_front() {
257 let handle = tasks.spawn(async move {
258 let _permit = request.permit;
259 request.forward.await;
260 });
261 running.insert(handle.id(), stream.clone());
262 }
263 if queue.is_empty() { pending.remove(&stream); }
264 }
265 }
266 request = requests.recv(), if !closed => {
267 match request {
268 Some(request) => {
269 let stream = request.permit.stream.clone();
270 if running.values().any(|active| active == &stream) {
271 pending.entry(stream).or_default().push_back(request);
272 } else {
273 let handle = tasks.spawn(async move {
274 let _permit = request.permit;
275 request.forward.await;
276 });
277 running.insert(handle.id(), stream);
278 }
279 }
280 None => closed = true,
281 }
282 }
283 }
284 if closed && tasks.is_empty() {
285 break;
286 }
287 }
288}
289
290pub struct SessionManagerShutdown {
296 pub(super) signal: Option<oneshot::Sender<()>>,
297 pub(super) task: Option<tokio::task::JoinHandle<()>>,
298}
299
300impl SessionManagerShutdown {
301 pub async fn shutdown(mut self) -> Result<()> {
302 if let Some(signal) = self.signal.take() {
303 let _ = signal.send(());
304 }
305 if let Some(task) = self.task.take() {
306 task.await.context("session manager shutdown task failed")?;
307 }
308 Ok(())
309 }
310}
311
312impl Drop for SessionManagerShutdown {
313 fn drop(&mut self) {
314 if let Some(signal) = self.signal.take() {
315 let _ = signal.send(());
316 }
317 if let Some(task) = self.task.take() {
318 task.abort();
319 }
320 }
321}
322
323#[derive(Clone)]
324pub(crate) struct CoalescedUpdateSender {
325 pub(super) cpu: watch::Sender<SessionCpuTable>,
326 producers: Arc<Mutex<BTreeMap<String, Arc<()>>>>,
327 producer: Option<Arc<UpdateProducer>>,
328 pub(super) delegation: Option<DelegationSender>,
329 pub(super) observer: Option<Arc<DelegationPublisher>>,
330 pub(super) mailbox: Arc<Mutex<UpdateMailbox>>,
331 pub(super) wake: mpsc::Sender<()>,
332}
333
334struct UpdateProducer {
335 cpu: watch::Sender<SessionCpuTable>,
336 registry: Arc<Mutex<BTreeMap<String, Arc<()>>>>,
337 session_id: String,
338 identity: Arc<()>,
339}
340
341impl Drop for UpdateProducer {
342 fn drop(&mut self) {
343 let mut registry = self
344 .registry
345 .lock()
346 .expect("session producer registry poisoned");
347 if registry
348 .get(&self.session_id)
349 .is_some_and(|current| Arc::ptr_eq(current, &self.identity))
350 {
351 registry.remove(&self.session_id);
352 self.cpu
353 .send_if_modified(|table| table.remove(&self.session_id).is_some());
354 }
355 }
356}
357
358pub struct SessionManagerUpdates {
361 pub(super) mailbox: Arc<Mutex<UpdateMailbox>>,
362 pub(super) wake: mpsc::Receiver<()>,
363 delivered_work: Option<crate::upgrade::Work>,
366}
367
368pub(super) struct PendingUpdate {
369 update: SessionManagerUpdate,
370 work: Option<crate::upgrade::Work>,
371}
372
373#[derive(Default)]
374pub(super) struct UpdateMailbox {
375 pub(super) pending: BTreeMap<String, PendingUpdate>,
376 ready: VecDeque<String>,
377}
378
379impl UpdateMailbox {
380 fn enqueue(&mut self, pending: PendingUpdate) {
381 let session_id = pending.update.session_id.clone();
382 if self.pending.insert(session_id.clone(), pending).is_none() {
385 self.ready.push_back(session_id);
386 }
387 }
388
389 fn remove(&mut self, session_id: &str) {
390 if self.pending.remove(session_id).is_some() {
391 self.ready.retain(|queued| queued != session_id);
392 }
393 }
394
395 fn pop(&mut self) -> Option<PendingUpdate> {
396 let session_id = self.ready.pop_front()?;
397 Some(
398 self.pending
399 .remove(&session_id)
400 .expect("queued session update"),
401 )
402 }
403}
404
405impl CoalescedUpdateSender {
406 pub(super) fn for_actor(&self, session_id: &str) -> Self {
407 assert!(
408 self.producer.is_none(),
409 "only the manager registers producers"
410 );
411 let identity = Arc::new(());
412 let mut registry = self
413 .producers
414 .lock()
415 .expect("session producer registry poisoned");
416 registry.insert(session_id.to_owned(), identity.clone());
417 self.cpu
418 .send_if_modified(|table| table.remove(session_id).is_some());
419 self.mailbox
421 .lock()
422 .expect("session update coalescer poisoned")
423 .remove(session_id);
424 let mut sender = self.clone();
425 sender.producer = Some(Arc::new(UpdateProducer {
426 cpu: self.cpu.clone(),
427 registry: self.producers.clone(),
428 session_id: session_id.to_owned(),
429 identity,
430 }));
431 sender
432 }
433
434 pub(super) fn publish_cpu(
435 &self,
436 session_id: &str,
437 value: Option<mj_client::runtime_feed::SessionCpuView>,
438 ) {
439 let registry = self
440 .producers
441 .lock()
442 .expect("session producer registry poisoned");
443 if let Some(producer) = &self.producer
444 && !registry
445 .get(session_id)
446 .is_some_and(|current| Arc::ptr_eq(current, &producer.identity))
447 {
448 return;
449 }
450 self.cpu.send_if_modified(|table| {
451 if table.get(session_id) == value.as_ref() {
452 return false;
453 }
454 match value {
455 Some(value) => {
456 table.insert(session_id.to_owned(), value);
457 }
458 None => {
459 table.remove(session_id);
460 }
461 }
462 true
463 });
464 }
465
466 pub(crate) fn send(&self, update: SessionManagerUpdate) {
467 let registry = self
468 .producers
469 .lock()
470 .expect("session producer registry poisoned");
471 if let Some(producer) = &self.producer {
472 assert_eq!(producer.session_id, update.session_id);
473 if !registry
474 .get(&update.session_id)
475 .is_some_and(|current| Arc::ptr_eq(current, &producer.identity))
476 {
477 return;
478 }
479 }
480 if let Some(observer) = &self.observer {
481 observer.publish(&update.view);
482 }
483 if self.wake.is_closed() {
484 return;
485 }
486 self.mailbox
487 .lock()
488 .expect("session update coalescer poisoned")
489 .enqueue(PendingUpdate {
490 update,
491 work: crate::upgrade::activity("session update").ok(),
492 });
493 let _ = self.wake.try_send(());
494 }
495}
496
497impl SessionManagerUpdates {
498 pub(super) fn pop_pending(&mut self) -> Option<SessionManagerUpdate> {
499 self.delivered_work = None;
500 let pending = self
501 .mailbox
502 .lock()
503 .expect("session update coalescer poisoned")
504 .pop()?;
505 self.delivered_work = pending.work;
506 Some(pending.update)
507 }
508
509 pub async fn recv(&mut self) -> Option<SessionManagerUpdate> {
510 loop {
511 if let Some(update) = self.pop_pending() {
512 return Some(update);
513 }
514 self.wake.recv().await?;
515 }
516 }
517
518 pub fn try_recv(
519 &mut self,
520 ) -> std::result::Result<SessionManagerUpdate, mpsc::error::TryRecvError> {
521 if let Some(update) = self.pop_pending() {
522 return Ok(update);
523 }
524 self.wake.try_recv()?;
525 self.pop_pending().ok_or(mpsc::error::TryRecvError::Empty)
526 }
527}
528
529pub(crate) fn coalesced_update_channel() -> (CoalescedUpdateSender, SessionManagerUpdates) {
530 let mailbox = Arc::new(Mutex::new(UpdateMailbox::default()));
531 let (wake_tx, wake_rx) = mpsc::channel(1);
532 (
533 CoalescedUpdateSender {
534 cpu: watch::channel(SessionCpuTable::new()).0,
535 producers: Default::default(),
536 producer: None,
537 delegation: None,
538 observer: None,
539 mailbox: mailbox.clone(),
540 wake: wake_tx,
541 },
542 SessionManagerUpdates {
543 mailbox,
544 wake: wake_rx,
545 delivered_work: None,
546 },
547 )
548}