mirror of
https://github.com/pchuan98/codex.git
synced 2026-07-01 00:31:56 +08:00
Add WebRTC transport to realtime start (#16960)
Adds WebRTC startup to the experimental app-server `thread/realtime/start` method with an optional transport enum. The websocket path remains the default; WebRTC offers create the realtime session through the shared start flow and emit the answer SDP via `thread/realtime/sdp`. --------- Co-authored-by: Codex <noreply@openai.com>
This commit is contained in:
committed by
GitHub
Unverified
parent
6c36e7d688
commit
fb3dcfde1d
@@ -39,6 +39,8 @@ use codex_api::MemoriesClient as ApiMemoriesClient;
|
||||
use codex_api::MemorySummarizeInput as ApiMemorySummarizeInput;
|
||||
use codex_api::MemorySummarizeOutput as ApiMemorySummarizeOutput;
|
||||
use codex_api::RawMemory as ApiRawMemory;
|
||||
use codex_api::RealtimeCallClient as ApiRealtimeCallClient;
|
||||
use codex_api::RealtimeSessionConfig;
|
||||
use codex_api::Reasoning;
|
||||
use codex_api::RequestTelemetry;
|
||||
use codex_api::ReqwestTransport;
|
||||
@@ -443,6 +445,30 @@ impl ModelClient {
|
||||
.map_err(map_api_error)
|
||||
}
|
||||
|
||||
pub async fn create_realtime_call(
|
||||
&self,
|
||||
sdp: String,
|
||||
session_config: RealtimeSessionConfig,
|
||||
) -> Result<String> {
|
||||
self.create_realtime_call_with_headers(sdp, session_config, ApiHeaderMap::new())
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn create_realtime_call_with_headers(
|
||||
&self,
|
||||
sdp: String,
|
||||
session_config: RealtimeSessionConfig,
|
||||
extra_headers: ApiHeaderMap,
|
||||
) -> Result<String> {
|
||||
let client_setup = self.current_client_setup().await?;
|
||||
let transport = ReqwestTransport::new(build_reqwest_client());
|
||||
ApiRealtimeCallClient::new(transport, client_setup.api_provider, client_setup.api_auth)
|
||||
.create_with_session_and_headers(sdp, session_config, extra_headers)
|
||||
.await
|
||||
.map(|response| response.sdp)
|
||||
.map_err(map_api_error)
|
||||
}
|
||||
|
||||
/// Builds memory summaries for each provided normalized raw memory.
|
||||
///
|
||||
/// This is a unary call (no streaming) to `/v1/memories/trace_summarize`.
|
||||
|
||||
@@ -7037,6 +7037,7 @@ fn realtime_text_for_event(msg: &EventMsg) -> Option<String> {
|
||||
EventMsg::Error(_)
|
||||
| EventMsg::Warning(_)
|
||||
| EventMsg::RealtimeConversationStarted(_)
|
||||
| EventMsg::RealtimeConversationSdp(_)
|
||||
| EventMsg::RealtimeConversationRealtime(_)
|
||||
| EventMsg::RealtimeConversationClosed(_)
|
||||
| EventMsg::ModelReroute(_)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use crate::client::ModelClient;
|
||||
use crate::codex::Session;
|
||||
use crate::realtime_context::build_realtime_startup_context;
|
||||
use async_channel::Receiver;
|
||||
@@ -27,12 +28,14 @@ use codex_protocol::error::Result as CodexResult;
|
||||
use codex_protocol::protocol::CodexErrorInfo;
|
||||
use codex_protocol::protocol::ConversationAudioParams;
|
||||
use codex_protocol::protocol::ConversationStartParams;
|
||||
use codex_protocol::protocol::ConversationStartTransport;
|
||||
use codex_protocol::protocol::ConversationTextParams;
|
||||
use codex_protocol::protocol::ErrorEvent;
|
||||
use codex_protocol::protocol::Event;
|
||||
use codex_protocol::protocol::EventMsg;
|
||||
use codex_protocol::protocol::RealtimeConversationClosedEvent;
|
||||
use codex_protocol::protocol::RealtimeConversationRealtimeEvent;
|
||||
use codex_protocol::protocol::RealtimeConversationSdpEvent;
|
||||
use codex_protocol::protocol::RealtimeConversationStartedEvent;
|
||||
use codex_protocol::protocol::RealtimeHandoffRequested;
|
||||
use http::HeaderMap;
|
||||
@@ -139,6 +142,24 @@ struct ConversationState {
|
||||
realtime_active: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
struct RealtimeStart {
|
||||
api_provider: ApiProvider,
|
||||
extra_headers: Option<HeaderMap>,
|
||||
session_config: RealtimeSessionConfig,
|
||||
model_client: ModelClient,
|
||||
sdp: Option<String>,
|
||||
}
|
||||
|
||||
struct RealtimeStartOutput {
|
||||
realtime_active: Arc<AtomicBool>,
|
||||
connection: RealtimeStartConnection,
|
||||
}
|
||||
|
||||
enum RealtimeStartConnection {
|
||||
Websocket { events_rx: Receiver<RealtimeEvent> },
|
||||
Webrtc { sdp: String },
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
impl RealtimeConversationManager {
|
||||
pub(crate) fn new() -> Self {
|
||||
@@ -154,12 +175,7 @@ impl RealtimeConversationManager {
|
||||
.and_then(|state| state.realtime_active.load(Ordering::Relaxed).then_some(()))
|
||||
}
|
||||
|
||||
pub(crate) async fn start(
|
||||
&self,
|
||||
api_provider: ApiProvider,
|
||||
extra_headers: Option<HeaderMap>,
|
||||
session_config: RealtimeSessionConfig,
|
||||
) -> CodexResult<(Receiver<RealtimeEvent>, Arc<AtomicBool>)> {
|
||||
async fn start(&self, start: RealtimeStart) -> CodexResult<RealtimeStartOutput> {
|
||||
let previous_state = {
|
||||
let mut guard = self.state.lock().await;
|
||||
guard.take()
|
||||
@@ -167,10 +183,37 @@ impl RealtimeConversationManager {
|
||||
if let Some(state) = previous_state {
|
||||
stop_conversation_state(state, RealtimeFanoutTaskStop::Abort).await;
|
||||
}
|
||||
|
||||
self.start_inner(start).await
|
||||
}
|
||||
|
||||
async fn start_inner(&self, start: RealtimeStart) -> CodexResult<RealtimeStartOutput> {
|
||||
let RealtimeStart {
|
||||
api_provider,
|
||||
extra_headers,
|
||||
session_config,
|
||||
model_client,
|
||||
sdp,
|
||||
} = start;
|
||||
let session_kind = match session_config.event_parser {
|
||||
RealtimeEventParser::V1 => RealtimeSessionKind::V1,
|
||||
RealtimeEventParser::RealtimeV2 => RealtimeSessionKind::V2,
|
||||
};
|
||||
let realtime_active = Arc::new(AtomicBool::new(true));
|
||||
|
||||
if let Some(sdp) = sdp {
|
||||
let sdp = model_client
|
||||
.create_realtime_call_with_headers(
|
||||
sdp,
|
||||
session_config,
|
||||
extra_headers.unwrap_or_default(),
|
||||
)
|
||||
.await?;
|
||||
return Ok(RealtimeStartOutput {
|
||||
realtime_active,
|
||||
connection: RealtimeStartConnection::Webrtc { sdp },
|
||||
});
|
||||
}
|
||||
|
||||
let client = RealtimeWebsocketClient::new(api_provider);
|
||||
let connection = client
|
||||
@@ -216,7 +259,10 @@ impl RealtimeConversationManager {
|
||||
fanout_task: None,
|
||||
realtime_active: Arc::clone(&realtime_active),
|
||||
});
|
||||
Ok((events_rx, realtime_active))
|
||||
Ok(RealtimeStartOutput {
|
||||
realtime_active,
|
||||
connection: RealtimeStartConnection::Websocket { events_rx },
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn register_fanout_task(
|
||||
@@ -447,6 +493,7 @@ struct PreparedRealtimeConversationStart {
|
||||
requested_session_id: Option<String>,
|
||||
version: RealtimeWsVersion,
|
||||
session_config: RealtimeSessionConfig,
|
||||
transport: ConversationStartTransport,
|
||||
}
|
||||
|
||||
async fn prepare_realtime_start(
|
||||
@@ -460,16 +507,54 @@ async fn prepare_realtime_start(
|
||||
.auth_manager()
|
||||
.unwrap_or_else(|| Arc::clone(&sess.services.auth_manager));
|
||||
let auth = auth_manager.auth().await;
|
||||
let realtime_api_key = realtime_api_key(auth.as_ref(), &provider)?;
|
||||
let mut api_provider = provider.to_api_provider(Some(AuthMode::ApiKey))?;
|
||||
let config = sess.get_config().await;
|
||||
let transport = params
|
||||
.transport
|
||||
.unwrap_or(ConversationStartTransport::Websocket);
|
||||
let mut api_provider = if matches!(transport, ConversationStartTransport::Websocket) {
|
||||
provider.to_api_provider(Some(AuthMode::ApiKey))?
|
||||
} else {
|
||||
provider.to_api_provider(auth.as_ref().map(CodexAuth::auth_mode))?
|
||||
};
|
||||
if let Some(realtime_ws_base_url) = &config.experimental_realtime_ws_base_url {
|
||||
api_provider.base_url = realtime_ws_base_url.clone();
|
||||
}
|
||||
let version = config.realtime.version;
|
||||
let session_config =
|
||||
build_realtime_session_config(sess, params.prompt, params.session_id).await?;
|
||||
let requested_session_id = session_config.session_id.clone();
|
||||
let extra_headers = match transport {
|
||||
ConversationStartTransport::Websocket => {
|
||||
let realtime_api_key = realtime_api_key(auth.as_ref(), &provider)?;
|
||||
realtime_request_headers(
|
||||
requested_session_id.as_deref(),
|
||||
Some(realtime_api_key.as_str()),
|
||||
)?
|
||||
}
|
||||
ConversationStartTransport::Webrtc { .. } => {
|
||||
realtime_request_headers(requested_session_id.as_deref(), /*api_key*/ None)?
|
||||
}
|
||||
};
|
||||
Ok(PreparedRealtimeConversationStart {
|
||||
api_provider,
|
||||
extra_headers,
|
||||
requested_session_id,
|
||||
version,
|
||||
session_config,
|
||||
transport,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn build_realtime_session_config(
|
||||
sess: &Arc<Session>,
|
||||
prompt: String,
|
||||
session_id: Option<String>,
|
||||
) -> CodexResult<RealtimeSessionConfig> {
|
||||
let config = sess.get_config().await;
|
||||
let prompt = config
|
||||
.experimental_realtime_ws_backend_prompt
|
||||
.clone()
|
||||
.unwrap_or(params.prompt);
|
||||
.unwrap_or(prompt);
|
||||
let startup_context = match config.experimental_realtime_ws_startup_context.clone() {
|
||||
Some(startup_context) => startup_context,
|
||||
None => {
|
||||
@@ -484,8 +569,7 @@ async fn prepare_realtime_start(
|
||||
format!("{prompt}\n\n{startup_context}")
|
||||
};
|
||||
let model = config.experimental_realtime_ws_model.clone();
|
||||
let version = config.realtime.version;
|
||||
let event_parser = match version {
|
||||
let event_parser = match config.realtime.version {
|
||||
RealtimeWsVersion::V1 => RealtimeEventParser::V1,
|
||||
RealtimeWsVersion::V2 => RealtimeEventParser::RealtimeV2,
|
||||
};
|
||||
@@ -493,22 +577,12 @@ async fn prepare_realtime_start(
|
||||
RealtimeWsMode::Conversational => RealtimeSessionMode::Conversational,
|
||||
RealtimeWsMode::Transcription => RealtimeSessionMode::Transcription,
|
||||
};
|
||||
let requested_session_id = params.session_id.or(Some(sess.conversation_id.to_string()));
|
||||
let session_config = RealtimeSessionConfig {
|
||||
Ok(RealtimeSessionConfig {
|
||||
instructions: prompt,
|
||||
model,
|
||||
session_id: requested_session_id.clone(),
|
||||
session_id: Some(session_id.unwrap_or_else(|| sess.conversation_id.to_string())),
|
||||
event_parser,
|
||||
session_mode,
|
||||
};
|
||||
let extra_headers =
|
||||
realtime_request_headers(requested_session_id.as_deref(), realtime_api_key.as_str())?;
|
||||
Ok(PreparedRealtimeConversationStart {
|
||||
api_provider,
|
||||
extra_headers,
|
||||
requested_session_id,
|
||||
version,
|
||||
session_config,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -523,12 +597,21 @@ async fn handle_start_inner(
|
||||
requested_session_id,
|
||||
version,
|
||||
session_config,
|
||||
transport,
|
||||
} = prepared_start;
|
||||
info!("starting realtime conversation");
|
||||
let (events_rx, realtime_active) = sess
|
||||
.conversation
|
||||
.start(api_provider, extra_headers, session_config)
|
||||
.await?;
|
||||
let sdp = match transport {
|
||||
ConversationStartTransport::Websocket => None,
|
||||
ConversationStartTransport::Webrtc { sdp } => Some(sdp),
|
||||
};
|
||||
let start = RealtimeStart {
|
||||
api_provider,
|
||||
extra_headers,
|
||||
session_config,
|
||||
model_client: sess.services.model_client.clone(),
|
||||
sdp,
|
||||
};
|
||||
let start_output = sess.conversation.start(start).await?;
|
||||
|
||||
info!("realtime conversation started");
|
||||
|
||||
@@ -541,6 +624,29 @@ async fn handle_start_inner(
|
||||
})
|
||||
.await;
|
||||
|
||||
let RealtimeStartOutput {
|
||||
realtime_active,
|
||||
connection,
|
||||
} = start_output;
|
||||
let events_rx = match connection {
|
||||
RealtimeStartConnection::Websocket { events_rx } => events_rx,
|
||||
RealtimeStartConnection::Webrtc { sdp } => {
|
||||
sess.send_event_raw(Event {
|
||||
id: sub_id.to_string(),
|
||||
msg: EventMsg::RealtimeConversationSdp(RealtimeConversationSdpEvent { sdp }),
|
||||
})
|
||||
.await;
|
||||
sess.conversation.finish_if_active(&realtime_active).await;
|
||||
send_realtime_conversation_closed(
|
||||
sess,
|
||||
sub_id.to_string(),
|
||||
RealtimeConversationEnd::TransportClosed,
|
||||
)
|
||||
.await;
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
let sess_clone = Arc::clone(sess);
|
||||
let sub_id = sub_id.to_string();
|
||||
let fanout_realtime_active = Arc::clone(&realtime_active);
|
||||
@@ -660,7 +766,7 @@ fn realtime_api_key(auth: Option<&CodexAuth>, provider: &ModelProviderInfo) -> C
|
||||
|
||||
fn realtime_request_headers(
|
||||
session_id: Option<&str>,
|
||||
api_key: &str,
|
||||
api_key: Option<&str>,
|
||||
) -> CodexResult<Option<HeaderMap>> {
|
||||
let mut headers = HeaderMap::new();
|
||||
|
||||
@@ -670,10 +776,12 @@ fn realtime_request_headers(
|
||||
headers.insert("x-session-id", session_id);
|
||||
}
|
||||
|
||||
let auth_value = HeaderValue::from_str(&format!("Bearer {api_key}")).map_err(|err| {
|
||||
CodexErr::InvalidRequest(format!("invalid realtime api key header: {err}"))
|
||||
})?;
|
||||
headers.insert(AUTHORIZATION, auth_value);
|
||||
if let Some(api_key) = api_key {
|
||||
let auth_value = HeaderValue::from_str(&format!("Bearer {api_key}")).map_err(|err| {
|
||||
CodexErr::InvalidRequest(format!("invalid realtime api key header: {err}"))
|
||||
})?;
|
||||
headers.insert(AUTHORIZATION, auth_value);
|
||||
}
|
||||
|
||||
Ok(Some(headers))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user