use futures::{Stream, stream};
use rig_http::http_client::Error;
use std::{
future::{Future, poll_fn},
pin::pin,
sync::LazyLock,
};
use tokio::runtime::{Handle, Runtime};
static RUNTIME: LazyLock<Result<Runtime, RuntimeUnavailable>> = LazyLock::new(|| {
tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.thread_name("rig-reqwest")
.enable_all()
.build()
.map_err(|err| RuntimeUnavailable(err.to_string()))
});
#[derive(Debug, Clone, thiserror::Error)]
#[error("rig-reqwest: failed to start the fallback tokio runtime: {0}")]
struct RuntimeUnavailable(String);
fn context() -> Result<Handle, Error> {
match Handle::try_current() {
Ok(handle) => Ok(handle),
Err(_) => RUNTIME
.as_ref()
.map(|runtime| runtime.handle().clone())
.map_err(|error| Error::instance(error.clone())),
}
}
pub(crate) fn bind<F: Future>(future: F) -> Result<impl Future<Output = F::Output>, Error> {
let handle = context()?;
Ok(async move {
let mut future = pin!(future);
poll_fn(|cx| {
let _entered = handle.enter();
future.as_mut().poll(cx)
})
.await
})
}
pub(crate) fn bind_stream<S: Stream>(stream: S) -> Result<impl Stream<Item = S::Item>, Error> {
let handle = context()?;
let mut stream = Box::pin(stream);
Ok(stream::poll_fn(move |cx| {
let _entered = handle.enter();
stream.as_mut().poll_next(cx)
}))
}