use rpi_ai::provider::SimpleStreamOptions;
use rpi_ai::types::Context;
use rpi_ai::{AssistantMessageEventStream, Model};
use std::sync::{Arc, OnceLock};
pub type StreamFn = Arc<
dyn Fn(&Model, &Context, &SimpleStreamOptions) -> AssistantMessageEventStream + Send + Sync,
>;
pub fn stream_fn<F>(f: F) -> StreamFn
where
F: Fn(&Model, &Context, &SimpleStreamOptions) -> AssistantMessageEventStream
+ Send
+ Sync
+ 'static,
{
Arc::new(f)
}
static DEFAULT_STREAM_FN: OnceLock<StreamFn> = OnceLock::new();
pub fn set_default_stream_fn(stream_fn: Option<StreamFn>) {
if let Some(f) = stream_fn {
let _ = DEFAULT_STREAM_FN.set(f);
}
}
pub fn try_get_default_stream_fn() -> Option<&'static StreamFn> {
DEFAULT_STREAM_FN.get()
}
pub fn get_default_stream_fn() -> Result<StreamFn, crate::AgentError> {
DEFAULT_STREAM_FN
.get()
.map(|f| Arc::clone(f))
.ok_or_else(|| {
crate::AgentError::State(
"no default stream fn configured — pass one to AgentBuilder or call \
set_default_stream_fn"
.into(),
)
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stream_fn_is_arc_clone() {
let f = stream_fn(|_, _, _| {
let (_prod, stream) = rpi_ai::event_stream::create_assistant_message_event_stream();
stream
});
let _clone: StreamFn = Arc::clone(&f);
}
}