mod session;
mod tools;
use std::future::IntoFuture;
use std::pin::Pin;
use std::sync::Arc;
use std::task::Context;
use std::task::Poll;
use std::time::Duration;
use ferrin_spec::JsonValue;
use ferrin_spec::RealtimeModelRef;
use ferrin_spec::realtime_model::ClientSecretOptions;
use ferrin_tool::ToolSet;
use futures_core::Stream;
use tokio::sync::mpsc;
use tokio::task::JoinSet;
use tokio_util::sync::CancellationToken;
pub use ferrin_spec::realtime_model::ClientSecret;
pub use ferrin_spec::realtime_model::ConversationItem;
pub use ferrin_spec::realtime_model::ConversationRole;
pub use ferrin_spec::realtime_model::Modality;
pub use ferrin_spec::realtime_model::RealtimeClientEvent;
pub use ferrin_spec::realtime_model::RealtimeServerEvent;
pub use ferrin_spec::realtime_model::RealtimeSessionConfig;
pub use ferrin_spec::realtime_model::RealtimeToolDefinition;
pub use ferrin_spec::realtime_model::ResponseCreateOptions;
pub use ferrin_spec::realtime_model::TranscriptionConfig;
pub use ferrin_spec::realtime_model::TurnDetection;
pub use ferrin_spec::realtime_model::TurnDetectionKind;
pub use session::RealtimeHandle;
pub use tools::realtime_tool_definitions;
use crate::error::Error;
use crate::registry::ProviderRegistry;
use crate::registry::default::resolve_model;
const DEFAULT_EVENT_BUFFER: usize = 256;
const CLOSE_TIMEOUT: Duration = Duration::from_secs(5);
pub type RealtimeEvent = Result<RealtimeServerEvent, Error>;
#[must_use]
pub fn realtime_session(model: impl Into<RealtimeModelRef>) -> RealtimeSessionBuilder {
RealtimeSessionBuilder {
model: model.into(),
client_secret: None,
expires_after_seconds: None,
config: RealtimeSessionConfig::default(),
tools: ToolSet::new(),
tools_context: None,
cancellation: CancellationToken::new(),
event_buffer: DEFAULT_EVENT_BUFFER,
}
}
#[derive(Debug)]
pub struct RealtimeSessionBuilder {
model: RealtimeModelRef,
client_secret: Option<ClientSecret>,
expires_after_seconds: Option<u64>,
config: RealtimeSessionConfig,
tools: ToolSet,
tools_context: Option<JsonValue>,
cancellation: CancellationToken,
event_buffer: usize,
}
impl RealtimeSessionBuilder {
#[must_use]
pub fn client_secret(mut self, secret: ClientSecret) -> Self {
self.client_secret = Some(secret);
self
}
#[must_use]
pub fn expires_after_seconds(mut self, seconds: u64) -> Self {
self.expires_after_seconds = Some(seconds);
self
}
#[must_use]
pub fn config(mut self, config: RealtimeSessionConfig) -> Self {
self.config = config;
self
}
#[must_use]
pub fn instructions(mut self, instructions: impl Into<String>) -> Self {
self.config.instructions = Some(instructions.into());
self
}
#[must_use]
pub fn voice(mut self, voice: impl Into<String>) -> Self {
self.config.voice = Some(voice.into());
self
}
#[must_use]
pub fn tools(mut self, tools: ToolSet) -> Self {
self.tools = tools;
self
}
#[must_use]
pub fn tools_context(mut self, context: JsonValue) -> Self {
self.tools_context = Some(context);
self
}
#[must_use]
pub fn cancellation(mut self, cancellation: CancellationToken) -> Self {
self.cancellation = cancellation;
self
}
#[must_use]
pub fn event_buffer(mut self, capacity: usize) -> Self {
self.event_buffer = capacity.max(1);
self
}
pub async fn connect(self) -> Result<RealtimeSession, Error> {
let model = resolve_model(&self.model, ProviderRegistry::realtime_model)?;
let mut config = self.config;
let definitions =
realtime_tool_definitions(&self.tools, self.tools_context.as_ref()).await?;
config.tools.extend(definitions);
let secret = match self.client_secret {
Some(secret) => secret,
None => model
.do_create_client_secret(ClientSecretOptions {
expires_after_seconds: self.expires_after_seconds,
session_config: Some(config.clone()),
})
.await
.map_err(Error::from)?,
};
let (events_tx, events_rx) = mpsc::channel(self.event_buffer);
let mut tasks = JoinSet::new();
let handle = session::start(
session::StartOptions {
model,
secret,
config,
tools: Arc::new(self.tools),
tools_context: self.tools_context,
cancellation: self.cancellation,
events: events_tx,
},
&mut tasks,
)
.await?;
Ok(RealtimeSession {
handle,
events: events_rx,
tasks,
})
}
}
impl IntoFuture for RealtimeSessionBuilder {
type Output = Result<RealtimeSession, Error>;
type IntoFuture = futures_util::future::BoxFuture<'static, Self::Output>;
fn into_future(self) -> Self::IntoFuture {
Box::pin(self.connect())
}
}
#[derive(Debug)]
pub struct RealtimeSession {
handle: RealtimeHandle,
events: mpsc::Receiver<RealtimeEvent>,
tasks: JoinSet<()>,
}
impl RealtimeSession {
#[must_use]
pub fn handle(&self) -> RealtimeHandle {
self.handle.clone()
}
pub async fn send(&self, event: RealtimeClientEvent) -> Result<(), Error> {
self.handle.send(event).await
}
pub async fn send_text(&self, text: impl Into<String>) -> Result<(), Error> {
self.handle.send_text(text).await
}
pub async fn add_tool_output(&self, call_id: &str, output: &JsonValue) -> Result<(), Error> {
self.handle.add_tool_output(call_id, output).await
}
pub async fn next_event(&mut self) -> Option<RealtimeEvent> {
self.events.recv().await
}
pub fn events(&mut self) -> impl Stream<Item = RealtimeEvent> + Send + '_ {
self
}
#[must_use]
pub fn is_closed(&self) -> bool {
self.handle.is_closed()
}
pub async fn close(mut self) -> Result<(), Error> {
self.handle.close();
let started = tokio::time::Instant::now();
let deadline = started + CLOSE_TIMEOUT;
loop {
match tokio::time::timeout_at(deadline, self.tasks.join_next()).await {
Ok(None) => return Ok(()),
Ok(Some(_)) => {}
Err(_) => {
self.tasks.abort_all();
return Err(Error::Timeout {
scope: crate::timeout::TimeoutScope::Total,
elapsed: started.elapsed(),
});
}
}
}
}
}
impl Stream for RealtimeSession {
type Item = RealtimeEvent;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.events.poll_recv(cx)
}
}