mirror of
https://github.com/openai/codex.git
synced 2026-09-17 12:23:33 +00:00
Switch realtime TUI to WebRTC transport
This commit is contained in:
@@ -20,7 +20,9 @@ tokio-tungstenite = { workspace = true }
|
||||
tungstenite = { workspace = true }
|
||||
tracing = { workspace = true }
|
||||
eventsource-stream = { workspace = true }
|
||||
libwebrtc = "0.3.26"
|
||||
regex-lite = { workspace = true }
|
||||
reqwest = { workspace = true, features = ["json", "multipart"] }
|
||||
tokio-util = { workspace = true, features = ["codec"] }
|
||||
url = { workspace = true }
|
||||
|
||||
@@ -30,7 +32,6 @@ assert_matches = { workspace = true }
|
||||
pretty_assertions = { workspace = true }
|
||||
tokio-test = { workspace = true }
|
||||
wiremock = { workspace = true }
|
||||
reqwest = { workspace = true }
|
||||
|
||||
[lints]
|
||||
workspace = true
|
||||
|
||||
5
codex-rs/codex-api/build.rs
Normal file
5
codex-rs/codex-api/build.rs
Normal file
@@ -0,0 +1,5 @@
|
||||
fn main() {
|
||||
if std::env::var("CARGO_CFG_TARGET_OS").as_deref() == Ok("macos") {
|
||||
println!("cargo:rustc-link-arg=-ObjC");
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,8 @@
|
||||
pub mod compact;
|
||||
pub mod memories;
|
||||
pub mod models;
|
||||
pub mod realtime_websocket;
|
||||
pub mod realtime_webrtc;
|
||||
mod realtime_websocket;
|
||||
pub mod responses;
|
||||
pub mod responses_websocket;
|
||||
mod session;
|
||||
|
||||
544
codex-rs/codex-api/src/endpoint/realtime_webrtc/mod.rs
Normal file
544
codex-rs/codex-api/src/endpoint/realtime_webrtc/mod.rs
Normal file
@@ -0,0 +1,544 @@
|
||||
use crate::endpoint::realtime_websocket::parse_realtime_event;
|
||||
use crate::error::ApiError;
|
||||
use crate::provider::Provider;
|
||||
use codex_protocol::protocol::RealtimeEvent;
|
||||
use http::HeaderMap;
|
||||
use libwebrtc::MediaType;
|
||||
use libwebrtc::data_channel::DataChannel;
|
||||
use libwebrtc::data_channel::DataChannelInit;
|
||||
use libwebrtc::data_channel::DataChannelState;
|
||||
use libwebrtc::peer_connection::OfferOptions;
|
||||
use libwebrtc::peer_connection::PeerConnection;
|
||||
use libwebrtc::peer_connection_factory::PeerConnectionFactory;
|
||||
use libwebrtc::peer_connection_factory::RtcConfiguration;
|
||||
use libwebrtc::peer_connection_factory::native::PeerConnectionFactoryExt;
|
||||
use libwebrtc::rtp_transceiver::RtpTransceiverDirection;
|
||||
use libwebrtc::rtp_transceiver::RtpTransceiverInit;
|
||||
use libwebrtc::session_description::SdpType;
|
||||
use libwebrtc::session_description::SessionDescription;
|
||||
use reqwest::Client;
|
||||
use reqwest::multipart::Form;
|
||||
use serde_json::json;
|
||||
use std::sync::Arc;
|
||||
use std::sync::Mutex as StdMutex;
|
||||
use std::sync::atomic::AtomicBool;
|
||||
use std::sync::atomic::Ordering;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::sync::oneshot;
|
||||
use tokio::time::Duration;
|
||||
use tracing::debug;
|
||||
use tracing::info;
|
||||
use tracing::warn;
|
||||
use url::Url;
|
||||
|
||||
const REALTIME_CALLS_PATH: &str = "/v1/realtime/calls";
|
||||
const REALTIME_DATA_CHANNEL_LABEL: &str = "oai-events";
|
||||
const REALTIME_VOICE: &str = "marin";
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct RealtimeSessionConfig {
|
||||
pub instructions: String,
|
||||
pub model: Option<String>,
|
||||
pub session_id: Option<String>,
|
||||
}
|
||||
|
||||
pub struct RealtimeWebrtcConnection {
|
||||
writer: RealtimeWebrtcWriter,
|
||||
events: RealtimeWebrtcEvents,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct RealtimeWebrtcWriter {
|
||||
peer_connection: PeerConnection,
|
||||
data_channel: DataChannel,
|
||||
is_closed: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct RealtimeWebrtcEvents {
|
||||
rx_event: Arc<Mutex<mpsc::UnboundedReceiver<RealtimeEvent>>>,
|
||||
is_closed: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
pub struct RealtimeWebrtcClient {
|
||||
provider: Provider,
|
||||
}
|
||||
|
||||
impl RealtimeWebrtcConnection {
|
||||
pub async fn send_conversation_item_create(&self, text: String) -> Result<(), ApiError> {
|
||||
self.writer.send_conversation_item_create(text).await
|
||||
}
|
||||
|
||||
pub async fn send_conversation_handoff_append(
|
||||
&self,
|
||||
handoff_id: String,
|
||||
output_text: String,
|
||||
) -> Result<(), ApiError> {
|
||||
self.writer
|
||||
.send_conversation_handoff_append(handoff_id, output_text)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn send_function_call_output(
|
||||
&self,
|
||||
call_id: String,
|
||||
output_text: String,
|
||||
) -> Result<(), ApiError> {
|
||||
self.writer
|
||||
.send_function_call_output(call_id, output_text)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn send_response_create(&self) -> Result<(), ApiError> {
|
||||
self.writer.send_response_create().await
|
||||
}
|
||||
|
||||
pub async fn close(&self) -> Result<(), ApiError> {
|
||||
self.writer.close().await
|
||||
}
|
||||
|
||||
pub async fn next_event(&self) -> Result<Option<RealtimeEvent>, ApiError> {
|
||||
self.events.next_event().await
|
||||
}
|
||||
|
||||
pub fn writer(&self) -> RealtimeWebrtcWriter {
|
||||
self.writer.clone()
|
||||
}
|
||||
|
||||
pub fn events(&self) -> RealtimeWebrtcEvents {
|
||||
self.events.clone()
|
||||
}
|
||||
|
||||
fn new(
|
||||
peer_connection: PeerConnection,
|
||||
data_channel: DataChannel,
|
||||
rx_event: mpsc::UnboundedReceiver<RealtimeEvent>,
|
||||
is_closed: Arc<AtomicBool>,
|
||||
) -> Self {
|
||||
Self {
|
||||
writer: RealtimeWebrtcWriter {
|
||||
peer_connection,
|
||||
data_channel,
|
||||
is_closed: Arc::clone(&is_closed),
|
||||
},
|
||||
events: RealtimeWebrtcEvents {
|
||||
rx_event: Arc::new(Mutex::new(rx_event)),
|
||||
is_closed,
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl RealtimeWebrtcWriter {
|
||||
pub async fn send_conversation_item_create(&self, text: String) -> Result<(), ApiError> {
|
||||
self.send_json(json!({
|
||||
"type": "conversation.item.create",
|
||||
"item": {
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{
|
||||
"type": "input_text",
|
||||
"text": text,
|
||||
}],
|
||||
},
|
||||
}))
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn send_conversation_handoff_append(
|
||||
&self,
|
||||
_handoff_id: String,
|
||||
output_text: String,
|
||||
) -> Result<(), ApiError> {
|
||||
self.send_json(json!({
|
||||
"type": "conversation.item.create",
|
||||
"item": {
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{
|
||||
"type": "output_text",
|
||||
"text": output_text,
|
||||
}],
|
||||
},
|
||||
}))
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn send_function_call_output(
|
||||
&self,
|
||||
call_id: String,
|
||||
output_text: String,
|
||||
) -> Result<(), ApiError> {
|
||||
self.send_json(json!({
|
||||
"type": "conversation.item.create",
|
||||
"item": {
|
||||
"type": "function_call_output",
|
||||
"call_id": call_id,
|
||||
"output": json!({
|
||||
"content": output_text,
|
||||
}).to_string(),
|
||||
},
|
||||
}))
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn send_response_create(&self) -> Result<(), ApiError> {
|
||||
self.send_json(json!({ "type": "response.create" })).await
|
||||
}
|
||||
|
||||
pub async fn close(&self) -> Result<(), ApiError> {
|
||||
if self.is_closed.swap(true, Ordering::SeqCst) {
|
||||
return Ok(());
|
||||
}
|
||||
self.data_channel.close();
|
||||
self.peer_connection.close();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn send_json(&self, payload: serde_json::Value) -> Result<(), ApiError> {
|
||||
if self.is_closed.load(Ordering::SeqCst) {
|
||||
return Err(ApiError::Stream(
|
||||
"realtime WebRTC connection is closed".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let serialized = serde_json::to_vec(&payload).map_err(|err| {
|
||||
ApiError::Stream(format!("failed to serialize realtime event: {err}"))
|
||||
})?;
|
||||
self.data_channel.send(&serialized, false).map_err(|err| {
|
||||
ApiError::Stream(format!("failed to send realtime data channel event: {err}"))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl RealtimeWebrtcEvents {
|
||||
pub async fn next_event(&self) -> Result<Option<RealtimeEvent>, ApiError> {
|
||||
if self.is_closed.load(Ordering::SeqCst) {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
match self.rx_event.lock().await.recv().await {
|
||||
Some(event) => Ok(Some(event)),
|
||||
None => {
|
||||
self.is_closed.store(true, Ordering::SeqCst);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl RealtimeWebrtcClient {
|
||||
pub fn new(provider: Provider) -> Self {
|
||||
Self { provider }
|
||||
}
|
||||
|
||||
pub async fn connect(
|
||||
&self,
|
||||
config: RealtimeSessionConfig,
|
||||
extra_headers: HeaderMap,
|
||||
default_headers: HeaderMap,
|
||||
) -> Result<RealtimeWebrtcConnection, ApiError> {
|
||||
info!("initializing realtime WebRTC peer connection");
|
||||
let factory = PeerConnectionFactory::with_platform_adm();
|
||||
let peer_connection = factory
|
||||
.create_peer_connection(RtcConfiguration::default())
|
||||
.map_err(|err| {
|
||||
ApiError::Stream(format!("failed to create WebRTC peer connection: {err}"))
|
||||
})?;
|
||||
|
||||
// Negotiate an audio m-line and attach a local mic track backed by the platform ADM.
|
||||
let audio_transceiver = peer_connection
|
||||
.add_transceiver_for_media(
|
||||
MediaType::Audio,
|
||||
RtpTransceiverInit {
|
||||
direction: RtpTransceiverDirection::SendRecv,
|
||||
stream_ids: vec!["realtime".to_string()],
|
||||
send_encodings: Vec::new(),
|
||||
},
|
||||
)
|
||||
.map_err(|err| ApiError::Stream(format!("failed to add audio transceiver: {err}")))?;
|
||||
|
||||
let local_audio_source = factory.create_audio_source();
|
||||
let local_audio_track = factory.create_audio_track("realtime-mic", local_audio_source);
|
||||
audio_transceiver
|
||||
.sender()
|
||||
.set_track(Some(local_audio_track.into()))
|
||||
.map_err(|err| {
|
||||
ApiError::Stream(format!("failed to attach ADM audio track to sender: {err}"))
|
||||
})?;
|
||||
|
||||
let data_channel = peer_connection
|
||||
.create_data_channel(REALTIME_DATA_CHANNEL_LABEL, DataChannelInit::default())
|
||||
.map_err(|err| {
|
||||
ApiError::Stream(format!("failed to create realtime data channel: {err}"))
|
||||
})?;
|
||||
|
||||
let (tx_event, rx_event) = mpsc::unbounded_channel();
|
||||
let is_closed = Arc::new(AtomicBool::new(false));
|
||||
let (tx_open, rx_open) = oneshot::channel::<()>();
|
||||
let tx_open = Arc::new(StdMutex::new(Some(tx_open)));
|
||||
|
||||
{
|
||||
let tx_event = tx_event.clone();
|
||||
data_channel.on_message(Some(Box::new(move |buffer| {
|
||||
if buffer.binary {
|
||||
debug!(
|
||||
payload_len = buffer.data.len(),
|
||||
"ignoring binary realtime data channel message"
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
let payload = match std::str::from_utf8(buffer.data) {
|
||||
Ok(payload) => payload,
|
||||
Err(err) => {
|
||||
debug!("received non-utf8 realtime data channel message: {err}");
|
||||
return;
|
||||
}
|
||||
};
|
||||
if let Some(event) = parse_realtime_event(payload)
|
||||
&& tx_event.send(event).is_err()
|
||||
{
|
||||
debug!("dropping realtime event because receiver closed");
|
||||
}
|
||||
})));
|
||||
}
|
||||
|
||||
{
|
||||
let is_closed = Arc::clone(&is_closed);
|
||||
let tx_open = Arc::clone(&tx_open);
|
||||
data_channel.on_state_change(Some(Box::new(move |state| match state {
|
||||
DataChannelState::Connecting => {}
|
||||
DataChannelState::Open => {
|
||||
if let Ok(mut tx_open) = tx_open.lock()
|
||||
&& let Some(tx_open) = tx_open.take()
|
||||
{
|
||||
let _ = tx_open.send(());
|
||||
}
|
||||
}
|
||||
DataChannelState::Closing | DataChannelState::Closed => {
|
||||
is_closed.store(true, Ordering::SeqCst);
|
||||
}
|
||||
})));
|
||||
}
|
||||
|
||||
let offer = peer_connection
|
||||
.create_offer(OfferOptions {
|
||||
ice_restart: false,
|
||||
offer_to_receive_audio: true,
|
||||
offer_to_receive_video: false,
|
||||
})
|
||||
.await
|
||||
.map_err(|err| ApiError::Stream(format!("failed to create WebRTC offer: {err}")))?;
|
||||
|
||||
peer_connection
|
||||
.set_local_description(offer.clone())
|
||||
.await
|
||||
.map_err(|err| ApiError::Stream(format!("failed to set local description: {err}")))?;
|
||||
|
||||
let url = realtime_calls_url(&self.provider.base_url)?;
|
||||
let headers = merge_request_headers(&self.provider.headers, extra_headers, default_headers);
|
||||
info!(url = %url, "posting realtime WebRTC offer");
|
||||
let http_client = Client::new();
|
||||
let mut request = http_client
|
||||
.post(url)
|
||||
.multipart(session_form(&config, &offer)?);
|
||||
for (name, value) in &headers {
|
||||
request = request.header(name, value);
|
||||
}
|
||||
|
||||
let response = request.send().await.map_err(|err| {
|
||||
ApiError::Stream(format!("failed to post realtime WebRTC offer: {err}"))
|
||||
})?;
|
||||
let status = response.status();
|
||||
let answer_sdp = response.text().await.map_err(|err| {
|
||||
ApiError::Stream(format!("failed to read realtime WebRTC answer body: {err}"))
|
||||
})?;
|
||||
if !status.is_success() {
|
||||
return Err(ApiError::Stream(format!(
|
||||
"realtime WebRTC offer failed with HTTP {status}: {answer_sdp}"
|
||||
)));
|
||||
}
|
||||
|
||||
let answer = SessionDescription::parse(&answer_sdp, SdpType::Answer)
|
||||
.map_err(|err| ApiError::Stream(format!("failed to parse WebRTC answer SDP: {err}")))?;
|
||||
peer_connection
|
||||
.set_remote_description(answer)
|
||||
.await
|
||||
.map_err(|err| ApiError::Stream(format!("failed to set remote description: {err}")))?;
|
||||
|
||||
if tokio::time::timeout(Duration::from_secs(10), rx_open)
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
warn!("timed out waiting for realtime data channel to open");
|
||||
}
|
||||
|
||||
Ok(RealtimeWebrtcConnection::new(
|
||||
peer_connection,
|
||||
data_channel,
|
||||
rx_event,
|
||||
is_closed,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
fn realtime_calls_url(base_url: &str) -> Result<Url, ApiError> {
|
||||
let mut url =
|
||||
Url::parse(base_url).map_err(|err| ApiError::Stream(format!("invalid base URL: {err}")))?;
|
||||
url.set_path(REALTIME_CALLS_PATH);
|
||||
Ok(url)
|
||||
}
|
||||
|
||||
fn session_form(
|
||||
config: &RealtimeSessionConfig,
|
||||
offer: &SessionDescription,
|
||||
) -> Result<Form, ApiError> {
|
||||
let session_json = serde_json::to_string(&session_payload(config))
|
||||
.map_err(|err| ApiError::Stream(format!("failed to serialize realtime session: {err}")))?;
|
||||
|
||||
Ok(Form::new()
|
||||
.text("sdp", offer.to_string())
|
||||
.text("session", session_json))
|
||||
}
|
||||
|
||||
fn session_payload(config: &RealtimeSessionConfig) -> serde_json::Value {
|
||||
let mut session = json!({
|
||||
"type": "realtime",
|
||||
"instructions": config.instructions,
|
||||
"output_modalities": ["audio"],
|
||||
"audio": {
|
||||
"output": {
|
||||
"voice": REALTIME_VOICE,
|
||||
},
|
||||
"input": {
|
||||
"turn_detection": {
|
||||
"type": "server_vad",
|
||||
"interrupt_response": true,
|
||||
"create_response": true,
|
||||
},
|
||||
},
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "codex",
|
||||
"description": "Delegate the user's request to Codex.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"input_transcript": {
|
||||
"type": "string",
|
||||
"description": "Transcript of the user's request.",
|
||||
},
|
||||
"send_immediately": {
|
||||
"type": "boolean",
|
||||
"description": "Whether Codex should receive the request immediately.",
|
||||
},
|
||||
},
|
||||
"required": ["input_transcript"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"name": "cancel_current_operation",
|
||||
"description": "Cancel the current Codex operation.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"required": [],
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"name": "turn_off_realtime_mode",
|
||||
"description": "Turn off realtime voice mode.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"required": [],
|
||||
},
|
||||
},
|
||||
],
|
||||
"tool_choice": "auto",
|
||||
});
|
||||
|
||||
if let Some(model) = &config.model {
|
||||
session["model"] = json!(model);
|
||||
}
|
||||
session
|
||||
}
|
||||
|
||||
fn merge_request_headers(
|
||||
provider_headers: &HeaderMap,
|
||||
extra_headers: HeaderMap,
|
||||
default_headers: HeaderMap,
|
||||
) -> HeaderMap {
|
||||
let mut headers = provider_headers.clone();
|
||||
headers.extend(extra_headers);
|
||||
for (name, value) in &default_headers {
|
||||
if let http::header::Entry::Vacant(entry) = headers.entry(name) {
|
||||
entry.insert(value.clone());
|
||||
}
|
||||
}
|
||||
headers
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use pretty_assertions::assert_eq;
|
||||
|
||||
#[test]
|
||||
fn realtime_calls_url_uses_webrtc_calls_path() {
|
||||
assert_eq!(
|
||||
realtime_calls_url("https://api.openai.com")
|
||||
.expect("url")
|
||||
.as_str(),
|
||||
"https://api.openai.com/v1/realtime/calls"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn session_form_contains_sdp_and_session_payload() {
|
||||
let offer = SessionDescription::parse(
|
||||
"v=0\r\no=- 0 0 IN IP4 127.0.0.1\r\ns=-\r\nt=0 0\r\n",
|
||||
SdpType::Offer,
|
||||
)
|
||||
.expect("offer");
|
||||
|
||||
let form = session_form(
|
||||
&RealtimeSessionConfig {
|
||||
instructions: "backend prompt".to_string(),
|
||||
model: Some("gpt-realtime-1.5".to_string()),
|
||||
session_id: Some("sess_123".to_string()),
|
||||
},
|
||||
&offer,
|
||||
)
|
||||
.expect("form");
|
||||
|
||||
let body = form.boundary().to_string();
|
||||
assert!(!body.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn session_payload_omits_session_id_from_body() {
|
||||
let payload = session_payload(&RealtimeSessionConfig {
|
||||
instructions: "backend prompt".to_string(),
|
||||
model: Some("gpt-realtime-1.5".to_string()),
|
||||
session_id: Some("sess_123".to_string()),
|
||||
});
|
||||
|
||||
assert_eq!(payload.get("id"), None);
|
||||
assert_eq!(payload.get("model"), Some(&json!("gpt-realtime-1.5")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn peer_connection_factory_smoke() {
|
||||
let factory = PeerConnectionFactory::default();
|
||||
let _pc = factory
|
||||
.create_peer_connection(RtcConfiguration::default())
|
||||
.expect("pc");
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,10 +1,3 @@
|
||||
pub mod methods;
|
||||
pub mod protocol;
|
||||
mod protocol;
|
||||
|
||||
pub use codex_protocol::protocol::RealtimeAudioFrame;
|
||||
pub use codex_protocol::protocol::RealtimeEvent;
|
||||
pub use methods::RealtimeWebsocketClient;
|
||||
pub use methods::RealtimeWebsocketConnection;
|
||||
pub use methods::RealtimeWebsocketEvents;
|
||||
pub use methods::RealtimeWebsocketWriter;
|
||||
pub use protocol::RealtimeSessionConfig;
|
||||
pub(crate) use protocol::parse_realtime_event;
|
||||
|
||||
@@ -1,150 +1,18 @@
|
||||
pub use codex_protocol::protocol::RealtimeAudioFrame;
|
||||
pub use codex_protocol::protocol::RealtimeCloseRequested;
|
||||
pub use codex_protocol::protocol::RealtimeEvent;
|
||||
pub use codex_protocol::protocol::RealtimeHandoffRequested;
|
||||
pub use codex_protocol::protocol::RealtimeInputAudioSpeechStarted;
|
||||
pub use codex_protocol::protocol::RealtimeInterruptRequested;
|
||||
pub use codex_protocol::protocol::RealtimeOutputAudioDelta;
|
||||
pub use codex_protocol::protocol::RealtimeResponseCancelled;
|
||||
pub use codex_protocol::protocol::RealtimeToolAction;
|
||||
pub use codex_protocol::protocol::RealtimeToolActionRequested;
|
||||
pub use codex_protocol::protocol::RealtimeTranscriptDelta;
|
||||
pub use codex_protocol::protocol::RealtimeTranscriptEntry;
|
||||
use codex_protocol::protocol::RealtimeCloseRequested;
|
||||
use codex_protocol::protocol::RealtimeEvent;
|
||||
use codex_protocol::protocol::RealtimeHandoffRequested;
|
||||
use codex_protocol::protocol::RealtimeInputAudioSpeechStarted;
|
||||
use codex_protocol::protocol::RealtimeInterruptRequested;
|
||||
use codex_protocol::protocol::RealtimeResponseCancelled;
|
||||
use codex_protocol::protocol::RealtimeToolAction;
|
||||
use codex_protocol::protocol::RealtimeToolActionRequested;
|
||||
use codex_protocol::protocol::RealtimeTranscriptDelta;
|
||||
use serde::Deserialize;
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
use std::collections::BTreeMap;
|
||||
use std::string::ToString;
|
||||
use tracing::debug;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct RealtimeSessionConfig {
|
||||
pub instructions: String,
|
||||
pub model: Option<String>,
|
||||
pub session_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub(super) enum RealtimeOutboundMessage {
|
||||
#[serde(rename = "input_audio_buffer.append")]
|
||||
InputAudioBufferAppend { audio: String },
|
||||
#[serde(rename = "response.create")]
|
||||
ResponseCreate,
|
||||
#[serde(rename = "conversation.item.truncate")]
|
||||
ConversationItemTruncate {
|
||||
item_id: String,
|
||||
content_index: u32,
|
||||
audio_end_ms: u32,
|
||||
},
|
||||
#[serde(rename = "session.update")]
|
||||
SessionUpdate { session: Box<SessionUpdateSession> },
|
||||
#[serde(rename = "conversation.item.create")]
|
||||
ConversationItemCreate { item: ConversationItem },
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionUpdateSession {
|
||||
#[serde(rename = "type")]
|
||||
pub(super) kind: String,
|
||||
pub(super) instructions: String,
|
||||
pub(super) output_modalities: Vec<String>,
|
||||
pub(super) audio: SessionAudio,
|
||||
pub(super) tools: Vec<SessionTool>,
|
||||
pub(super) tool_choice: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionAudio {
|
||||
pub(super) input: SessionAudioInput,
|
||||
pub(super) output: SessionAudioOutput,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionAudioInput {
|
||||
pub(super) format: SessionAudioFormat,
|
||||
pub(super) noise_reduction: SessionNoiseReduction,
|
||||
pub(super) turn_detection: SessionTurnDetection,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionAudioFormat {
|
||||
#[serde(rename = "type")]
|
||||
pub(super) kind: String,
|
||||
pub(super) rate: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionNoiseReduction {
|
||||
#[serde(rename = "type")]
|
||||
pub(super) kind: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionTurnDetection {
|
||||
#[serde(rename = "type")]
|
||||
pub(super) kind: String,
|
||||
pub(super) interrupt_response: bool,
|
||||
pub(super) create_response: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionAudioOutput {
|
||||
pub(super) format: SessionAudioOutputFormat,
|
||||
pub(super) voice: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionAudioOutputFormat {
|
||||
#[serde(rename = "type")]
|
||||
pub(super) kind: String,
|
||||
pub(super) rate: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionTool {
|
||||
#[serde(rename = "type")]
|
||||
pub(super) kind: String,
|
||||
pub(super) name: String,
|
||||
pub(super) description: String,
|
||||
pub(super) parameters: SessionToolParameters,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionToolParameters {
|
||||
#[serde(rename = "type")]
|
||||
pub(super) kind: String,
|
||||
pub(super) properties: BTreeMap<String, SessionToolProperty>,
|
||||
pub(super) required: Vec<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct SessionToolProperty {
|
||||
#[serde(rename = "type")]
|
||||
pub(super) kind: String,
|
||||
pub(super) description: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub(super) enum ConversationItem {
|
||||
#[serde(rename = "message")]
|
||||
Message {
|
||||
role: String,
|
||||
content: Vec<ConversationItemContent>,
|
||||
},
|
||||
#[serde(rename = "function_call_output")]
|
||||
FunctionCallOutput { call_id: String, output: String },
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(super) struct ConversationItemContent {
|
||||
#[serde(rename = "type")]
|
||||
pub(super) kind: String,
|
||||
pub(super) text: String,
|
||||
}
|
||||
|
||||
pub(super) fn parse_realtime_event(payload: &str) -> Option<RealtimeEvent> {
|
||||
pub(crate) fn parse_realtime_event(payload: &str) -> Option<RealtimeEvent> {
|
||||
let parsed: Value = match serde_json::from_str(payload) {
|
||||
Ok(msg) => msg,
|
||||
Err(err) => {
|
||||
@@ -160,6 +28,7 @@ pub(super) fn parse_realtime_event(payload: &str) -> Option<RealtimeEvent> {
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
match message_type {
|
||||
"session.created" | "session.updated" => {
|
||||
let session_id = parsed
|
||||
@@ -179,42 +48,6 @@ pub(super) fn parse_realtime_event(payload: &str) -> Option<RealtimeEvent> {
|
||||
instructions,
|
||||
})
|
||||
}
|
||||
|
||||
"conversation.output_audio.delta"
|
||||
| "response.output_audio.delta"
|
||||
| "response.audio.delta" => {
|
||||
let data = parsed
|
||||
.get("delta")
|
||||
.and_then(Value::as_str)
|
||||
.or_else(|| parsed.get("data").and_then(Value::as_str))
|
||||
.map(str::to_string)?;
|
||||
let sample_rate = parsed
|
||||
.get("sample_rate")
|
||||
.and_then(Value::as_u64)
|
||||
.and_then(|v| u32::try_from(v).ok())
|
||||
.unwrap_or(24_000);
|
||||
let num_channels = parsed
|
||||
.get("channels")
|
||||
.or_else(|| parsed.get("num_channels"))
|
||||
.and_then(Value::as_u64)
|
||||
.and_then(|v| u16::try_from(v).ok())
|
||||
.unwrap_or(1);
|
||||
Some(RealtimeEvent::AudioOut(RealtimeOutputAudioDelta {
|
||||
frame: RealtimeAudioFrame {
|
||||
data,
|
||||
sample_rate,
|
||||
num_channels,
|
||||
samples_per_channel: parsed
|
||||
.get("samples_per_channel")
|
||||
.and_then(Value::as_u64)
|
||||
.and_then(|v| u32::try_from(v).ok()),
|
||||
},
|
||||
item_id: parsed
|
||||
.get("item_id")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::to_string),
|
||||
}))
|
||||
}
|
||||
"input_audio_buffer.speech_started" => Some(RealtimeEvent::InputAudioSpeechStarted(
|
||||
RealtimeInputAudioSpeechStarted {
|
||||
item_id: parsed
|
||||
@@ -525,13 +358,6 @@ fn parse_tool_action_requested(parsed: &Value) -> Option<RealtimeToolActionReque
|
||||
});
|
||||
}
|
||||
|
||||
if let Some(function_call) = find_function_call(parsed, "compact_conversation") {
|
||||
return Some(RealtimeToolActionRequested {
|
||||
call_id: parse_call_id(function_call)?,
|
||||
action: RealtimeToolAction::CompactConversation,
|
||||
});
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
@@ -551,80 +377,39 @@ fn parse_response_cancelled(parsed: &Value) -> Option<RealtimeResponseCancelled>
|
||||
}
|
||||
|
||||
fn find_function_call<'a>(parsed: &'a Value, name: &str) -> Option<&'a Value> {
|
||||
parsed
|
||||
let output = parsed
|
||||
.get("response")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|response| response.get("output"))
|
||||
.and_then(Value::as_array)?
|
||||
.iter()
|
||||
.find(|item| {
|
||||
item.get("type").and_then(Value::as_str) == Some("function_call")
|
||||
&& item.get("name").and_then(Value::as_str) == Some(name)
|
||||
})
|
||||
.and_then(Value::as_array)?;
|
||||
output.iter().find(|item| {
|
||||
item.get("type").and_then(Value::as_str) == Some("function_call")
|
||||
&& item.get("name").and_then(Value::as_str) == Some(name)
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
struct ParsedHandoffArguments {
|
||||
input_transcript: String,
|
||||
send_immediately: bool,
|
||||
}
|
||||
|
||||
fn parse_handoff_arguments(arguments: &str) -> ParsedHandoffArguments {
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct HandoffArguments {
|
||||
#[derive(Debug, Deserialize, Default)]
|
||||
struct RawHandoffArguments {
|
||||
#[serde(default)]
|
||||
prompt: Option<String>,
|
||||
#[serde(default)]
|
||||
text: Option<String>,
|
||||
#[serde(default)]
|
||||
input: Option<String>,
|
||||
#[serde(default)]
|
||||
message: Option<String>,
|
||||
#[serde(default)]
|
||||
input_transcript: Option<String>,
|
||||
input_transcript: String,
|
||||
#[serde(default)]
|
||||
send_immediately: bool,
|
||||
#[serde(default)]
|
||||
messages: Vec<RealtimeTranscriptEntry>,
|
||||
}
|
||||
|
||||
let Some(parsed) = serde_json::from_str::<HandoffArguments>(arguments).ok() else {
|
||||
return ParsedHandoffArguments {
|
||||
serde_json::from_str::<RawHandoffArguments>(arguments)
|
||||
.map(|raw| ParsedHandoffArguments {
|
||||
input_transcript: raw.input_transcript,
|
||||
send_immediately: raw.send_immediately,
|
||||
})
|
||||
.unwrap_or_else(|_| ParsedHandoffArguments {
|
||||
input_transcript: arguments.to_string(),
|
||||
send_immediately: false,
|
||||
};
|
||||
};
|
||||
|
||||
for value in [
|
||||
parsed.prompt,
|
||||
parsed.text,
|
||||
parsed.input,
|
||||
parsed.message,
|
||||
parsed.input_transcript,
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
{
|
||||
if !value.is_empty() {
|
||||
return ParsedHandoffArguments {
|
||||
input_transcript: value,
|
||||
send_immediately: parsed.send_immediately,
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(message) = parsed
|
||||
.messages
|
||||
.into_iter()
|
||||
.find(|message| message.role == "user" && !message.text.is_empty())
|
||||
{
|
||||
return ParsedHandoffArguments {
|
||||
input_transcript: message.text,
|
||||
send_immediately: parsed.send_immediately,
|
||||
};
|
||||
}
|
||||
|
||||
ParsedHandoffArguments {
|
||||
input_transcript: String::new(),
|
||||
send_immediately: parsed.send_immediately,
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -27,9 +27,9 @@ pub use crate::common::create_text_param_for_request;
|
||||
pub use crate::endpoint::compact::CompactClient;
|
||||
pub use crate::endpoint::memories::MemoriesClient;
|
||||
pub use crate::endpoint::models::ModelsClient;
|
||||
pub use crate::endpoint::realtime_websocket::RealtimeSessionConfig;
|
||||
pub use crate::endpoint::realtime_websocket::RealtimeWebsocketClient;
|
||||
pub use crate::endpoint::realtime_websocket::RealtimeWebsocketConnection;
|
||||
pub use crate::endpoint::realtime_webrtc::RealtimeSessionConfig;
|
||||
pub use crate::endpoint::realtime_webrtc::RealtimeWebrtcClient;
|
||||
pub use crate::endpoint::realtime_webrtc::RealtimeWebrtcConnection;
|
||||
pub use crate::endpoint::responses::ResponsesClient;
|
||||
pub use crate::endpoint::responses::ResponsesOptions;
|
||||
pub use crate::endpoint::responses_websocket::ResponsesWebsocketClient;
|
||||
|
||||
@@ -1,482 +0,0 @@
|
||||
use std::collections::HashMap;
|
||||
use std::future::Future;
|
||||
use std::time::Duration;
|
||||
|
||||
use codex_api::RealtimeAudioFrame;
|
||||
use codex_api::RealtimeEvent;
|
||||
use codex_api::RealtimeSessionConfig;
|
||||
use codex_api::RealtimeWebsocketClient;
|
||||
use codex_api::provider::Provider;
|
||||
use codex_api::provider::RetryConfig;
|
||||
use codex_protocol::protocol::RealtimeOutputAudioDelta;
|
||||
use futures::SinkExt;
|
||||
use futures::StreamExt;
|
||||
use http::HeaderMap;
|
||||
use serde_json::Value;
|
||||
use serde_json::json;
|
||||
use tokio::net::TcpListener;
|
||||
use tokio_tungstenite::accept_async;
|
||||
use tokio_tungstenite::tungstenite::Message;
|
||||
|
||||
type RealtimeWsStream = tokio_tungstenite::WebSocketStream<tokio::net::TcpStream>;
|
||||
|
||||
async fn spawn_realtime_ws_server<Handler, Fut>(
|
||||
handler: Handler,
|
||||
) -> (String, tokio::task::JoinHandle<()>)
|
||||
where
|
||||
Handler: FnOnce(RealtimeWsStream) -> Fut + Send + 'static,
|
||||
Fut: Future<Output = ()> + Send + 'static,
|
||||
{
|
||||
let listener = match TcpListener::bind("127.0.0.1:0").await {
|
||||
Ok(listener) => listener,
|
||||
Err(err) => panic!("failed to bind test websocket listener: {err}"),
|
||||
};
|
||||
let addr = match listener.local_addr() {
|
||||
Ok(addr) => addr.to_string(),
|
||||
Err(err) => panic!("failed to read local websocket listener address: {err}"),
|
||||
};
|
||||
|
||||
let server = tokio::spawn(async move {
|
||||
let (stream, _) = match listener.accept().await {
|
||||
Ok(stream) => stream,
|
||||
Err(err) => panic!("failed to accept test websocket connection: {err}"),
|
||||
};
|
||||
let ws = match accept_async(stream).await {
|
||||
Ok(ws) => ws,
|
||||
Err(err) => panic!("failed to complete websocket handshake: {err}"),
|
||||
};
|
||||
handler(ws).await;
|
||||
});
|
||||
|
||||
(addr, server)
|
||||
}
|
||||
|
||||
fn test_provider(base_url: String) -> Provider {
|
||||
Provider {
|
||||
name: "test".to_string(),
|
||||
base_url,
|
||||
query_params: Some(HashMap::new()),
|
||||
headers: HeaderMap::new(),
|
||||
retry: RetryConfig {
|
||||
max_attempts: 1,
|
||||
base_delay: Duration::from_millis(1),
|
||||
retry_429: false,
|
||||
retry_5xx: false,
|
||||
retry_transport: false,
|
||||
},
|
||||
stream_idle_timeout: Duration::from_secs(5),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn realtime_ws_e2e_session_create_and_event_flow() {
|
||||
let (addr, server) = spawn_realtime_ws_server(|mut ws: RealtimeWsStream| async move {
|
||||
let first = ws
|
||||
.next()
|
||||
.await
|
||||
.expect("first msg")
|
||||
.expect("first msg ok")
|
||||
.into_text()
|
||||
.expect("text");
|
||||
let first_json: Value = serde_json::from_str(&first).expect("json");
|
||||
assert_eq!(first_json["type"], "session.update");
|
||||
assert_eq!(
|
||||
first_json["session"]["type"],
|
||||
Value::String("realtime".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["instructions"],
|
||||
Value::String("backend prompt".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["audio"]["input"]["format"]["type"],
|
||||
Value::String("audio/pcm".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["audio"]["input"]["format"]["rate"],
|
||||
Value::from(24_000)
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["audio"]["input"]["noise_reduction"]["type"],
|
||||
Value::String("near_field".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["audio"]["input"]["turn_detection"]["type"],
|
||||
Value::String("server_vad".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["audio"]["input"]["turn_detection"]["interrupt_response"],
|
||||
Value::Bool(true)
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["audio"]["input"]["turn_detection"]["create_response"],
|
||||
Value::Bool(true)
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["audio"]["output"]["format"]["type"],
|
||||
Value::String("audio/pcm".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["audio"]["output"]["format"]["rate"],
|
||||
Value::from(24_000)
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["audio"]["output"]["voice"],
|
||||
Value::String("marin".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tool_choice"],
|
||||
Value::String("auto".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"][0]["type"],
|
||||
Value::String("function".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"][0]["name"],
|
||||
Value::String("codex".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"][0]["parameters"]["properties"]["send_immediately"]
|
||||
["type"],
|
||||
Value::String("boolean".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"][1]["type"],
|
||||
Value::String("function".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"][1]["name"],
|
||||
Value::String("cancel_current_operation".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"][1]["parameters"]["type"],
|
||||
Value::String("object".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"][1]["parameters"]["required"],
|
||||
Value::Array(Vec::new())
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"][1]["parameters"]["properties"],
|
||||
json!({})
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"]
|
||||
.as_array()
|
||||
.expect("tools array")
|
||||
.iter()
|
||||
.map(|tool| tool["name"].as_str().expect("tool name"))
|
||||
.collect::<Vec<_>>(),
|
||||
vec![
|
||||
"codex",
|
||||
"cancel_current_operation",
|
||||
"turn_off_realtime_mode",
|
||||
"manage_message_queue",
|
||||
"manage_runtime_settings",
|
||||
"run_tui_command",
|
||||
]
|
||||
);
|
||||
assert_eq!(
|
||||
first_json["session"]["tools"][4]["parameters"]["properties"]["working_directory"]
|
||||
["type"],
|
||||
Value::String("string".to_string())
|
||||
);
|
||||
|
||||
ws.send(Message::Text(
|
||||
json!({
|
||||
"type": "session.updated",
|
||||
"session": {"id": "sess_mock", "instructions": "backend prompt"}
|
||||
})
|
||||
.to_string()
|
||||
.into(),
|
||||
))
|
||||
.await
|
||||
.expect("send session.updated");
|
||||
|
||||
let second = ws
|
||||
.next()
|
||||
.await
|
||||
.expect("second msg")
|
||||
.expect("second msg ok")
|
||||
.into_text()
|
||||
.expect("text");
|
||||
let second_json: Value = serde_json::from_str(&second).expect("json");
|
||||
assert_eq!(second_json["type"], "input_audio_buffer.append");
|
||||
|
||||
ws.send(Message::Text(
|
||||
json!({
|
||||
"type": "response.output_audio.delta",
|
||||
"delta": "AQID",
|
||||
"sample_rate": 48000,
|
||||
"channels": 1
|
||||
})
|
||||
.to_string()
|
||||
.into(),
|
||||
))
|
||||
.await
|
||||
.expect("send audio out");
|
||||
})
|
||||
.await;
|
||||
|
||||
let client = RealtimeWebsocketClient::new(test_provider(format!("http://{addr}")));
|
||||
let connection = client
|
||||
.connect(
|
||||
RealtimeSessionConfig {
|
||||
instructions: "backend prompt".to_string(),
|
||||
model: Some("realtime-test-model".to_string()),
|
||||
session_id: Some("conv_123".to_string()),
|
||||
},
|
||||
HeaderMap::new(),
|
||||
HeaderMap::new(),
|
||||
)
|
||||
.await
|
||||
.expect("connect");
|
||||
|
||||
let created = connection
|
||||
.next_event()
|
||||
.await
|
||||
.expect("next event")
|
||||
.expect("event");
|
||||
assert_eq!(
|
||||
created,
|
||||
RealtimeEvent::SessionUpdated {
|
||||
session_id: "sess_mock".to_string(),
|
||||
instructions: Some("backend prompt".to_string()),
|
||||
}
|
||||
);
|
||||
|
||||
connection
|
||||
.send_audio_frame(RealtimeAudioFrame {
|
||||
data: "AQID".to_string(),
|
||||
sample_rate: 48000,
|
||||
num_channels: 1,
|
||||
samples_per_channel: Some(960),
|
||||
})
|
||||
.await
|
||||
.expect("send audio");
|
||||
|
||||
let audio_event = connection
|
||||
.next_event()
|
||||
.await
|
||||
.expect("next event")
|
||||
.expect("event");
|
||||
assert_eq!(
|
||||
audio_event,
|
||||
RealtimeEvent::AudioOut(RealtimeOutputAudioDelta {
|
||||
frame: RealtimeAudioFrame {
|
||||
data: "AQID".to_string(),
|
||||
sample_rate: 48000,
|
||||
num_channels: 1,
|
||||
samples_per_channel: None,
|
||||
},
|
||||
item_id: None,
|
||||
})
|
||||
);
|
||||
|
||||
connection.close().await.expect("close");
|
||||
server.await.expect("server task");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn realtime_ws_e2e_send_while_next_event_waits() {
|
||||
let (addr, server) = spawn_realtime_ws_server(|mut ws: RealtimeWsStream| async move {
|
||||
let first = ws
|
||||
.next()
|
||||
.await
|
||||
.expect("first msg")
|
||||
.expect("first msg ok")
|
||||
.into_text()
|
||||
.expect("text");
|
||||
let first_json: Value = serde_json::from_str(&first).expect("json");
|
||||
assert_eq!(first_json["type"], "session.update");
|
||||
|
||||
let second = ws
|
||||
.next()
|
||||
.await
|
||||
.expect("second msg")
|
||||
.expect("second msg ok")
|
||||
.into_text()
|
||||
.expect("text");
|
||||
let second_json: Value = serde_json::from_str(&second).expect("json");
|
||||
assert_eq!(second_json["type"], "input_audio_buffer.append");
|
||||
|
||||
ws.send(Message::Text(
|
||||
json!({
|
||||
"type": "session.updated",
|
||||
"session": {"id": "sess_after_send", "instructions": "backend prompt"}
|
||||
})
|
||||
.to_string()
|
||||
.into(),
|
||||
))
|
||||
.await
|
||||
.expect("send session.updated");
|
||||
})
|
||||
.await;
|
||||
|
||||
let client = RealtimeWebsocketClient::new(test_provider(format!("http://{addr}")));
|
||||
let connection = client
|
||||
.connect(
|
||||
RealtimeSessionConfig {
|
||||
instructions: "backend prompt".to_string(),
|
||||
model: Some("realtime-test-model".to_string()),
|
||||
session_id: Some("conv_123".to_string()),
|
||||
},
|
||||
HeaderMap::new(),
|
||||
HeaderMap::new(),
|
||||
)
|
||||
.await
|
||||
.expect("connect");
|
||||
|
||||
let (send_result, next_result) = tokio::join!(
|
||||
async {
|
||||
tokio::time::timeout(
|
||||
Duration::from_millis(200),
|
||||
connection.send_audio_frame(RealtimeAudioFrame {
|
||||
data: "AQID".to_string(),
|
||||
sample_rate: 48000,
|
||||
num_channels: 1,
|
||||
samples_per_channel: Some(960),
|
||||
}),
|
||||
)
|
||||
.await
|
||||
},
|
||||
connection.next_event()
|
||||
);
|
||||
|
||||
send_result
|
||||
.expect("send should not block on next_event")
|
||||
.expect("send audio");
|
||||
let next_event = next_result.expect("next event").expect("event");
|
||||
assert_eq!(
|
||||
next_event,
|
||||
RealtimeEvent::SessionUpdated {
|
||||
session_id: "sess_after_send".to_string(),
|
||||
instructions: Some("backend prompt".to_string()),
|
||||
}
|
||||
);
|
||||
|
||||
connection.close().await.expect("close");
|
||||
server.await.expect("server task");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn realtime_ws_e2e_disconnected_emitted_once() {
|
||||
let (addr, server) = spawn_realtime_ws_server(|mut ws: RealtimeWsStream| async move {
|
||||
let first = ws
|
||||
.next()
|
||||
.await
|
||||
.expect("first msg")
|
||||
.expect("first msg ok")
|
||||
.into_text()
|
||||
.expect("text");
|
||||
let first_json: Value = serde_json::from_str(&first).expect("json");
|
||||
assert_eq!(first_json["type"], "session.update");
|
||||
|
||||
ws.send(Message::Close(None)).await.expect("send close");
|
||||
})
|
||||
.await;
|
||||
|
||||
let client = RealtimeWebsocketClient::new(test_provider(format!("http://{addr}")));
|
||||
let connection = client
|
||||
.connect(
|
||||
RealtimeSessionConfig {
|
||||
instructions: "backend prompt".to_string(),
|
||||
model: Some("realtime-test-model".to_string()),
|
||||
session_id: Some("conv_123".to_string()),
|
||||
},
|
||||
HeaderMap::new(),
|
||||
HeaderMap::new(),
|
||||
)
|
||||
.await
|
||||
.expect("connect");
|
||||
|
||||
let first = connection.next_event().await.expect("next event");
|
||||
assert_eq!(first, None);
|
||||
|
||||
let second = connection.next_event().await.expect("next event");
|
||||
assert_eq!(second, None);
|
||||
|
||||
server.await.expect("server task");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn realtime_ws_e2e_forwards_unknown_text_events() {
|
||||
let (addr, server) = spawn_realtime_ws_server(|mut ws: RealtimeWsStream| async move {
|
||||
let first = ws
|
||||
.next()
|
||||
.await
|
||||
.expect("first msg")
|
||||
.expect("first msg ok")
|
||||
.into_text()
|
||||
.expect("text");
|
||||
let first_json: Value = serde_json::from_str(&first).expect("json");
|
||||
assert_eq!(first_json["type"], "session.update");
|
||||
|
||||
ws.send(Message::Text(
|
||||
json!({
|
||||
"type": "response.created",
|
||||
"response": {"id": "resp_unknown"}
|
||||
})
|
||||
.to_string()
|
||||
.into(),
|
||||
))
|
||||
.await
|
||||
.expect("send unknown event");
|
||||
|
||||
ws.send(Message::Text(
|
||||
json!({
|
||||
"type": "session.updated",
|
||||
"session": {"id": "sess_after_unknown", "instructions": "backend prompt"}
|
||||
})
|
||||
.to_string()
|
||||
.into(),
|
||||
))
|
||||
.await
|
||||
.expect("send session.updated");
|
||||
})
|
||||
.await;
|
||||
|
||||
let client = RealtimeWebsocketClient::new(test_provider(format!("http://{addr}")));
|
||||
let connection = client
|
||||
.connect(
|
||||
RealtimeSessionConfig {
|
||||
instructions: "backend prompt".to_string(),
|
||||
model: Some("realtime-test-model".to_string()),
|
||||
session_id: Some("conv_123".to_string()),
|
||||
},
|
||||
HeaderMap::new(),
|
||||
HeaderMap::new(),
|
||||
)
|
||||
.await
|
||||
.expect("connect");
|
||||
|
||||
let first_event = connection
|
||||
.next_event()
|
||||
.await
|
||||
.expect("next event")
|
||||
.expect("event");
|
||||
assert_eq!(
|
||||
first_event,
|
||||
RealtimeEvent::ConversationItemAdded(json!({
|
||||
"type": "response.created",
|
||||
"response": {"id": "resp_unknown"}
|
||||
}))
|
||||
);
|
||||
|
||||
let second_event = connection
|
||||
.next_event()
|
||||
.await
|
||||
.expect("next event")
|
||||
.expect("event");
|
||||
assert_eq!(
|
||||
second_event,
|
||||
RealtimeEvent::SessionUpdated {
|
||||
session_id: "sess_after_unknown".to_string(),
|
||||
instructions: Some("backend prompt".to_string()),
|
||||
}
|
||||
);
|
||||
|
||||
connection.close().await.expect("close");
|
||||
server.await.expect("server task");
|
||||
}
|
||||
Reference in New Issue
Block a user