use meerkat_core::types::{ContentInput, HandlingMode};
use meerkat_mob::{
MemberTurnEventSender, MemberTurnHandle, MemberTurnOptions, MobError, MobHandle,
SpawnMemberSpec, SpawnResult,
};
use std::future::Future;
use crate::mob_handle_runtime::MobRuntimeError;
use super::UnifiedRuntime;
const MAX_CONCURRENT_SPAWN_MANY: usize = 1;
#[derive(Debug)]
pub struct MemberTurnAdmission {
pub session_id: String,
pub turn: MemberTurnHandle,
}
pub(crate) const DEFAULT_WAIT_READY_TIMEOUT: std::time::Duration =
std::time::Duration::from_mins(10);
pub(crate) fn is_ready_wait_timeout(err: &MobError) -> bool {
matches!(err, MobError::ReadyWaitTimedOut { .. })
}
impl UnifiedRuntime {
pub fn mob_handle(&self) -> MobHandle {
self.mob_runtime.handle()
}
pub fn mob_runtime(&self) -> &crate::mob_handle_runtime::MobRuntime {
&self.mob_runtime
}
pub async fn start_member_turn(
&self,
member_alias: &str,
content: ContentInput,
handling_mode: HandlingMode,
options: MemberTurnOptions,
event_tx: Option<MemberTurnEventSender>,
) -> Result<MemberTurnAdmission, MobRuntimeError> {
let identity = crate::member_comms_id::mob_member_id(member_alias);
let member = self.mob_handle().member(&identity).await?;
let turn = member
.start_turn(content, handling_mode, options, event_tx)
.await?;
let session_id = turn
.session_id()
.ok_or(MobRuntimeError::InvalidInput(
"completion-bearing member turns require a session-backed member",
))?
.to_string();
Ok(MemberTurnAdmission { session_id, turn })
}
pub async fn spawn(&self, mut spec: SpawnMemberSpec) -> Result<SpawnResult, MobRuntimeError> {
if let Some(labels) = spec.labels.as_ref() {
crate::member_comms_id::validate_raw_identity_labels(labels)
.map_err(|message| MobRuntimeError::InvalidConfig(message.to_string()))?;
}
let raw_reservation = crate::member_comms_id::reserve_raw_member_target(
self.identity_runtime(),
spec.identity.as_str(),
)
.await
.map_err(MobRuntimeError::InvalidConfig)?;
let member_id = raw_reservation.alias().to_string();
let profile = spec.role_name.to_string();
spec.identity = crate::member_comms_id::mob_member_id(member_id.as_str());
let spawn_result = Box::pin(self.mob_handle().spawn_spec(spec)).await;
drop(raw_reservation);
match spawn_result {
Ok(result) => {
if let Some(hook) = &self.post_spawn_hook {
hook(vec![member_id]).await;
}
Ok(result)
}
Err(err) => {
self.fire_error(super::types::ErrorEvent::SpawnFailure {
member_id,
profile,
error: format!("{err}"),
});
Err(err.into())
}
}
}
pub async fn spawn_many(
&self,
mut specs: Vec<SpawnMemberSpec>,
) -> Result<Vec<SpawnResult>, MobRuntimeError> {
let requested_member_ids = specs
.iter()
.map(|spec| spec.identity.to_string())
.collect::<Vec<_>>();
for spec in &specs {
if let Some(labels) = spec.labels.as_ref() {
crate::member_comms_id::validate_raw_identity_labels(labels)
.map_err(|message| MobRuntimeError::InvalidConfig(message.to_string()))?;
}
}
let raw_reservation = if specs.is_empty() {
None
} else {
Some(
crate::member_comms_id::reserve_raw_member_targets(
self.identity_runtime(),
requested_member_ids.iter().map(String::as_str),
)
.await
.map_err(MobRuntimeError::InvalidConfig)?,
)
};
let member_ids = raw_reservation
.as_ref()
.map(|reservation| reservation.aliases().to_vec())
.unwrap_or(requested_member_ids);
for (spec, member_id) in specs.iter_mut().zip(&member_ids) {
spec.identity = crate::member_comms_id::mob_member_id(member_id);
}
let handle = self.mob_handle();
let refs = try_join_in_batches(specs, MAX_CONCURRENT_SPAWN_MANY, |spec| {
let handle = handle.clone();
async move { Box::pin(handle.spawn_spec(spec)).await }
})
.await
.map_err(MobRuntimeError::from)?;
drop(raw_reservation);
if !member_ids.is_empty()
&& let Some(hook) = &self.post_spawn_hook
{
hook(member_ids).await;
}
Ok(refs)
}
}
async fn try_join_in_batches<I, F, T, E, Build>(
items: Vec<I>,
batch_size: usize,
mut build: Build,
) -> Result<Vec<T>, E>
where
F: Future<Output = Result<T, E>>,
Build: FnMut(I) -> F,
{
let batch_size = batch_size.max(1);
let mut results = Vec::with_capacity(items.len());
let mut iter = items.into_iter();
loop {
let batch: Vec<I> = iter.by_ref().take(batch_size).collect();
if batch.is_empty() {
break;
}
let futures = batch.into_iter().map(&mut build);
let mut batch_results = futures::future::try_join_all(futures).await?;
results.append(&mut batch_results);
tokio::task::yield_now().await;
}
Ok(results)
}
#[cfg(test)]
mod tests {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use super::{is_ready_wait_timeout, try_join_in_batches};
use meerkat_mob::MobError;
#[tokio::test]
async fn spawn_many_batch_size_stays_serial_until_upstream_backpressure_exists() {
assert_eq!(super::MAX_CONCURRENT_SPAWN_MANY, 1);
}
#[test]
fn ready_wait_timeout_is_classified_as_envelope_not_error() {
let timed_out = MobError::ReadyWaitTimedOut {
pending_member_ids: vec![],
};
assert!(is_ready_wait_timeout(&timed_out));
let display = timed_out.to_string().to_lowercase();
assert!(
!display.contains("timeout"),
"old substring check would have missed this timeout"
);
assert!(display.contains("timed out"));
assert!(!is_ready_wait_timeout(&MobError::KickoffWaitTimedOut {
pending_member_ids: vec![],
}));
}
#[tokio::test]
async fn try_join_in_batches_can_run_serially() {
let active = Arc::new(AtomicUsize::new(0));
let max_active = Arc::new(AtomicUsize::new(0));
let items: Vec<usize> = (0..25).collect();
let results = try_join_in_batches(items.clone(), 1, |item| {
let active = active.clone();
let max_active = max_active.clone();
async move {
let current = active.fetch_add(1, Ordering::SeqCst) + 1;
max_active.fetch_max(current, Ordering::SeqCst);
tokio::task::yield_now().await;
active.fetch_sub(1, Ordering::SeqCst);
Ok::<_, ()>(item)
}
})
.await;
assert_eq!(results, Ok(items));
assert_eq!(max_active.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn try_join_in_batches_limits_concurrent_work_and_preserves_order() {
let active = Arc::new(AtomicUsize::new(0));
let max_active = Arc::new(AtomicUsize::new(0));
let items: Vec<usize> = (0..75).collect();
let results = try_join_in_batches(items.clone(), 16, |item| {
let active = active.clone();
let max_active = max_active.clone();
async move {
let current = active.fetch_add(1, Ordering::SeqCst) + 1;
max_active.fetch_max(current, Ordering::SeqCst);
tokio::task::yield_now().await;
active.fetch_sub(1, Ordering::SeqCst);
Ok::<_, ()>(item)
}
})
.await;
assert_eq!(results, Ok(items));
assert!(max_active.load(Ordering::SeqCst) <= 16);
}
#[tokio::test]
async fn try_join_in_batches_stops_before_starting_later_batches_after_error() {
let started = Arc::new(AtomicUsize::new(0));
let items: Vec<usize> = (0..40).collect();
let result = try_join_in_batches(items, 16, |item| {
let started = started.clone();
async move {
started.fetch_add(1, Ordering::SeqCst);
tokio::task::yield_now().await;
if item == 20 { Err(item) } else { Ok(item) }
}
})
.await;
assert_eq!(result, Err(20));
assert_eq!(started.load(Ordering::SeqCst), 32);
}
}