1use std::fmt;
35use std::net::SocketAddr;
36use std::sync::{Arc, Mutex, PoisonError};
37use std::time::Duration;
38
39use base64::Engine as _;
40use base64::engine::general_purpose::STANDARD as BASE64;
41use futures_util::{SinkExt, StreamExt};
42use serde_json::{Value, json};
43use tokio::net::{TcpListener, TcpStream};
44use tokio::sync::{Notify, mpsc, oneshot};
45use tokio::task::{JoinHandle, JoinSet};
46use tokio_tungstenite::accept_hdr_async;
47use tokio_tungstenite::tungstenite::handshake::server::{ErrorResponse, Request, Response};
48use tokio_tungstenite::tungstenite::http::StatusCode;
49use tokio_tungstenite::tungstenite::protocol::{CloseFrame, frame::coding::CloseCode};
50use tokio_tungstenite::tungstenite::{Bytes as WsBytes, Message, Utf8Bytes};
51use tokio_util::sync::CancellationToken;
52
53pub const FRAME_BYTES: usize = 160;
56
57pub const F_SILENCE: [u8; FRAME_BYTES] = [0xFF; FRAME_BYTES];
59
60pub const F_SILENCE_BASE64: &str = "/////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////////w==";
66
67pub const F_RAMP_BASE64: &str = "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8gISIjJCUmJygpKissLS4vMDEyMzQ1Njc4OTo7PD0+P0BBQkNERUZHSElKS0xNTk9QUVJTVFVWV1hZWltcXV5fYGFiY2RlZmdoaWprbG1ub3BxcnN0dXZ3eHl6e3x9fn+AgYKDhIWGh4iJiouMjY6PkJGSk5SVlpeYmZqbnJ2enw==";
72
73pub const FIXTURE_BEARER: &str = "fixture-bearer";
78
79pub const OBSERVATION_BOUND: Duration = Duration::from_secs(10);
85
86#[must_use]
97pub fn tone_frame(index: usize) -> [u8; FRAME_BYTES] {
98 let mut frame = [0u8; FRAME_BYTES];
99 let base = index.wrapping_mul(FRAME_BYTES);
100 for (offset, byte) in frame.iter_mut().enumerate() {
101 *byte = u8::try_from(base.wrapping_add(offset) % 256).unwrap_or_default();
102 }
103 frame
104}
105
106#[must_use]
109pub fn tone_bytes(frames: usize) -> Vec<u8> {
110 (0..frames).flat_map(tone_frame).collect()
111}
112
113#[derive(Debug, Clone, Copy, PartialEq, Eq)]
117pub enum UpgradeOutcome {
118 Accepted,
120 Refused(u16),
123}
124
125#[derive(Debug, Clone)]
132pub struct Upgrade {
133 pub target: String,
135 pub authorization: Option<String>,
137 pub header_names: Vec<String>,
140 pub outcome: UpgradeOutcome,
142}
143
144#[derive(Debug, Clone)]
151pub enum ClientEvent {
152 SessionUpdate(Value),
155 Append {
157 audio: Vec<u8>,
159 },
160 Cancel,
162 Outside {
165 event_type: String,
167 },
168 Unreadable {
171 reason: String,
173 },
174}
175
176#[derive(Debug, Clone, Default)]
182pub struct Record {
183 pub upgrades: Vec<Upgrade>,
185 pub client_events: Vec<ClientEvent>,
187 pub appended_audio: Vec<u8>,
190 pub pings: usize,
193 pub deltas_sent: usize,
195 pub deltas_suppressed: usize,
198 pub sessions_ended: usize,
200}
201
202impl Record {
203 #[must_use]
205 pub fn appends(&self) -> usize {
206 self.client_events
207 .iter()
208 .filter(|event| matches!(event, ClientEvent::Append { .. }))
209 .count()
210 }
211
212 #[must_use]
214 pub fn cancels(&self) -> usize {
215 self.client_events
216 .iter()
217 .filter(|event| matches!(event, ClientEvent::Cancel))
218 .count()
219 }
220
221 #[must_use]
223 pub fn session_updates(&self) -> Vec<&Value> {
224 self.client_events
225 .iter()
226 .filter_map(|event| match event {
227 ClientEvent::SessionUpdate(update) => Some(update),
228 _ => None,
229 })
230 .collect()
231 }
232
233 #[must_use]
237 pub fn events_outside_the_client_subset(&self) -> Vec<String> {
238 self.client_events
239 .iter()
240 .filter_map(|event| match event {
241 ClientEvent::Outside { event_type } => Some(event_type.clone()),
242 ClientEvent::Unreadable { reason } => Some(format!("unreadable: {reason}")),
243 _ => None,
244 })
245 .collect()
246 }
247
248 #[must_use]
250 pub fn accepted(&self) -> usize {
251 self.outcomes(UpgradeOutcome::Accepted)
252 }
253
254 #[must_use]
256 pub fn refused(&self) -> usize {
257 self.upgrades
258 .iter()
259 .filter(|upgrade| matches!(upgrade.outcome, UpgradeOutcome::Refused(_)))
260 .count()
261 }
262
263 fn outcomes(&self, outcome: UpgradeOutcome) -> usize {
264 self.upgrades
265 .iter()
266 .filter(|upgrade| upgrade.outcome == outcome)
267 .count()
268 }
269}
270
271#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
279pub enum Withhold {
280 #[default]
282 Nothing,
283 SessionCreated,
285 SessionUpdated,
287}
288
289#[derive(Debug, Clone, Copy, PartialEq, Eq)]
294pub enum StallPoint {
295 Upgrade,
298 Session,
302}
303
304#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
306pub enum CancelPolicy {
307 #[default]
311 Truncate,
312 KeepStreaming,
316}
317
318#[derive(Debug, Clone)]
323pub enum Malformed {
324 NotJson,
326 NoType,
328 Binary,
330 DeltaNotBase64 {
332 response: String,
334 },
335 DeltaMissing {
337 response: String,
339 },
340 AudioDoneWithoutResponseId,
342}
343
344#[derive(Debug, Clone, Copy, PartialEq, Eq)]
346pub enum Emission {
347 Sent,
349 SuppressedByCancel,
352}
353
354#[derive(Debug, thiserror::Error)]
356#[non_exhaustive]
357pub enum PeerError {
358 #[error("the stand-in peer could not bind a loopback listener: {0}")]
360 Bind(String),
361 #[error("the stand-in peer has no connected session")]
363 NoSession,
364 #[error("the stand-in peer's session ended before the directive was performed")]
366 SessionEnded,
367 #[error("the stand-in peer did not observe {what} within {OBSERVATION_BOUND:?}")]
369 NotObserved {
370 what: String,
372 },
373}
374
375#[derive(Debug, Clone)]
379pub struct PeerConfig {
380 bearer: String,
381 withhold: Withhold,
382 stall: Option<StallPoint>,
383 cancel: CancelPolicy,
384}
385
386impl Default for PeerConfig {
387 fn default() -> Self {
388 Self::new()
389 }
390}
391
392impl PeerConfig {
393 #[must_use]
395 pub fn new() -> Self {
396 Self {
397 bearer: FIXTURE_BEARER.to_owned(),
398 withhold: Withhold::Nothing,
399 stall: None,
400 cancel: CancelPolicy::Truncate,
401 }
402 }
403
404 #[must_use]
406 pub fn expecting_bearer(mut self, bearer: &str) -> Self {
407 bearer.clone_into(&mut self.bearer);
408 self
409 }
410
411 #[must_use]
413 pub fn withholding(mut self, withhold: Withhold) -> Self {
414 self.withhold = withhold;
415 self
416 }
417
418 #[must_use]
420 pub fn stalling_at(mut self, stall: StallPoint) -> Self {
421 self.stall = Some(stall);
422 self
423 }
424
425 #[must_use]
427 pub fn on_cancel(mut self, cancel: CancelPolicy) -> Self {
428 self.cancel = cancel;
429 self
430 }
431
432 pub async fn start(self) -> Result<RealtimePeer, PeerError> {
436 let listener = TcpListener::bind("127.0.0.1:0")
437 .await
438 .map_err(|error| PeerError::Bind(error.to_string()))?;
439 let addr = listener
440 .local_addr()
441 .map_err(|error| PeerError::Bind(error.to_string()))?;
442 let shared = Arc::new(Shared::default());
443 let shutdown = CancellationToken::new();
444 let accepting = tokio::spawn(accept(
445 listener,
446 self,
447 Arc::clone(&shared),
448 shutdown.clone(),
449 ));
450 Ok(RealtimePeer {
451 url: format!("ws://{addr}/v1/realtime"),
452 addr,
453 shared,
454 shutdown,
455 accepting: Some(accepting),
456 })
457 }
458}
459
460#[derive(Debug)]
465pub struct RealtimePeer {
466 url: String,
467 addr: SocketAddr,
468 shared: Arc<Shared>,
469 shutdown: CancellationToken,
470 accepting: Option<JoinHandle<()>>,
471}
472
473impl Drop for RealtimePeer {
474 fn drop(&mut self) {
475 self.shutdown.cancel();
476 }
477}
478
479impl RealtimePeer {
480 #[must_use]
482 pub fn url(&self) -> &str {
483 &self.url
484 }
485
486 #[must_use]
488 pub fn addr(&self) -> SocketAddr {
489 self.addr
490 }
491
492 #[must_use]
494 pub fn record(&self) -> Record {
495 self.shared.snapshot()
496 }
497
498 pub async fn observe<F>(&self, what: &str, condition: F) -> Result<Record, PeerError>
503 where
504 F: Fn(&Record) -> bool,
505 {
506 tokio::time::timeout(OBSERVATION_BOUND, async {
510 loop {
511 let changed = self.shared.changed.notified();
512 {
513 let record = self.shared.lock();
514 if condition(&record) {
515 return record.clone();
516 }
517 }
518 changed.await;
519 }
520 })
521 .await
522 .map_err(|_elapsed| PeerError::NotObserved {
523 what: what.to_owned(),
524 })
525 }
526
527 pub async fn await_upgrade(&self) -> Result<Record, PeerError> {
529 self.observe("an upgrade", |record| !record.upgrades.is_empty())
530 .await
531 }
532
533 pub async fn await_session_update(&self) -> Result<Record, PeerError> {
535 self.observe("a session.update", |record| {
536 !record.session_updates().is_empty()
537 })
538 .await
539 }
540
541 pub async fn await_appends(&self, count: usize) -> Result<Record, PeerError> {
543 self.observe(&format!("{count} appends"), move |record| {
544 record.appends() >= count
545 })
546 .await
547 }
548
549 pub async fn await_cancel(&self) -> Result<Record, PeerError> {
551 self.observe("a response.cancel", |record| record.cancels() > 0)
552 .await
553 }
554
555 pub async fn send_delta(&self, response: &str, audio: &[u8]) -> Result<Emission, PeerError> {
557 self.direct(Action::Delta {
558 response: response.to_owned(),
559 audio: audio.to_vec(),
560 })
561 .await
562 }
563
564 pub async fn speak_tone(&self, response: &str, frames: usize) -> Result<usize, PeerError> {
567 let mut sent = 0;
568 for frame in 0..frames {
569 if self.send_delta(response, &tone_frame(frame)).await? == Emission::SuppressedByCancel
570 {
571 break;
572 }
573 sent += 1;
574 }
575 Ok(sent)
576 }
577
578 pub async fn send_audio_done(&self, response: &str) -> Result<Emission, PeerError> {
580 self.direct(Action::Scripted(Scripted::AudioDone {
581 response: Some(response.to_owned()),
582 }))
583 .await
584 }
585
586 pub async fn send_response_done(
588 &self,
589 response: &str,
590 status: &str,
591 ) -> Result<Emission, PeerError> {
592 self.direct(Action::Scripted(Scripted::ResponseDone {
593 response: response.to_owned(),
594 status: status.to_owned(),
595 }))
596 .await
597 }
598
599 pub async fn send_speech_started(&self) -> Result<Emission, PeerError> {
601 self.direct(Action::Scripted(Scripted::SpeechStarted)).await
602 }
603
604 pub async fn send_error(&self, code: &str, message: &str) -> Result<Emission, PeerError> {
606 self.direct(Action::Scripted(Scripted::Error {
607 code: code.to_owned(),
608 message: message.to_owned(),
609 }))
610 .await
611 }
612
613 pub async fn send_unknown(&self, event_type: &str) -> Result<Emission, PeerError> {
615 self.direct(Action::Scripted(Scripted::Unknown {
616 event_type: event_type.to_owned(),
617 }))
618 .await
619 }
620
621 pub async fn send_malformed(&self, malformed: Malformed) -> Result<Emission, PeerError> {
623 self.direct(Action::Malformed(malformed)).await
624 }
625
626 pub async fn send_oversize(&self, bytes: usize) -> Result<Emission, PeerError> {
628 self.direct(Action::Scripted(Scripted::Oversize { bytes }))
629 .await
630 }
631
632 pub async fn close_normally(&self) -> Result<Emission, PeerError> {
634 self.direct(Action::Close { code: 1000 }).await
635 }
636
637 pub async fn reset(&self) -> Result<Emission, PeerError> {
639 self.direct(Action::Reset).await
640 }
641
642 pub async fn shutdown(mut self) {
647 self.shutdown.cancel();
648 if let Some(accepting) = self.accepting.take() {
649 let _joined = accepting.await;
650 }
651 }
652
653 async fn direct(&self, action: Action) -> Result<Emission, PeerError> {
654 let session = self.shared.session().ok_or(PeerError::NoSession)?;
655 let (done, performed) = oneshot::channel();
656 session
657 .send(Directive { action, done })
658 .await
659 .map_err(|_closed| PeerError::SessionEnded)?;
660 performed.await.map_err(|_closed| PeerError::SessionEnded)
661 }
662}
663
664#[derive(Debug, Default)]
668struct Shared {
669 record: Mutex<Record>,
670 changed: Notify,
671 session: Mutex<Option<(u64, mpsc::Sender<Directive>)>>,
672 generations: Mutex<u64>,
673}
674
675impl Shared {
676 fn lock(&self) -> std::sync::MutexGuard<'_, Record> {
677 self.record.lock().unwrap_or_else(PoisonError::into_inner)
678 }
679
680 fn update<F: FnOnce(&mut Record)>(&self, edit: F) {
681 edit(&mut self.lock());
682 self.changed.notify_waiters();
683 }
684
685 fn snapshot(&self) -> Record {
686 self.lock().clone()
687 }
688
689 fn register(&self, sender: mpsc::Sender<Directive>) -> u64 {
695 let mut generations = self
696 .generations
697 .lock()
698 .unwrap_or_else(PoisonError::into_inner);
699 *generations += 1;
700 let generation = *generations;
701 *self.session.lock().unwrap_or_else(PoisonError::into_inner) = Some((generation, sender));
702 generation
703 }
704
705 fn unregister(&self, generation: u64) {
706 let mut session = self.session.lock().unwrap_or_else(PoisonError::into_inner);
707 if session
708 .as_ref()
709 .is_some_and(|(open, _)| *open == generation)
710 {
711 *session = None;
712 }
713 }
714
715 fn session(&self) -> Option<mpsc::Sender<Directive>> {
716 self.session
717 .lock()
718 .unwrap_or_else(PoisonError::into_inner)
719 .as_ref()
720 .map(|(_, sender)| sender.clone())
721 }
722}
723
724#[derive(Debug)]
726struct Directive {
727 action: Action,
728 done: oneshot::Sender<Emission>,
729}
730
731#[derive(Debug)]
732enum Action {
733 Delta { response: String, audio: Vec<u8> },
735 Scripted(Scripted),
737 Malformed(Malformed),
739 Close { code: u16 },
741 Reset,
743}
744
745#[derive(Debug)]
746enum Scripted {
747 AudioDone { response: Option<String> },
748 ResponseDone { response: String, status: String },
749 SpeechStarted,
750 Error { code: String, message: String },
751 Unknown { event_type: String },
752 Oversize { bytes: usize },
753}
754
755type Socket = tokio_tungstenite::WebSocketStream<TcpStream>;
759
760#[derive(Debug, Default, PartialEq, Eq)]
767enum Cancelled {
768 #[default]
769 No,
770 Pending,
771 Response(String),
772}
773
774#[derive(Debug, Default)]
776struct Session {
777 in_flight: Option<String>,
779 cancelled: Cancelled,
781 events: u32,
782}
783
784impl Session {
785 fn next_event_id(&mut self) -> String {
786 self.events += 1;
787 format!("event_{:03}", self.events)
788 }
789
790 fn is_cancelled(&self, response: &str) -> bool {
792 match &self.cancelled {
793 Cancelled::No => false,
794 Cancelled::Pending => true,
795 Cancelled::Response(cancelled) => cancelled == response,
796 }
797 }
798}
799
800async fn accept(
801 listener: TcpListener,
802 config: PeerConfig,
803 shared: Arc<Shared>,
804 shutdown: CancellationToken,
805) {
806 let mut sessions = JoinSet::new();
807 loop {
808 tokio::select! {
809 () = shutdown.cancelled() => break,
810 accepted = listener.accept() => {
811 let Ok((stream, _from)) = accepted else { break };
812 sessions.spawn(serve(
813 stream,
814 config.clone(),
815 Arc::clone(&shared),
816 shutdown.child_token(),
817 ));
818 }
819 Some(_finished) = sessions.join_next(), if !sessions.is_empty() => {}
820 }
821 }
822 sessions.shutdown().await;
825}
826
827#[allow(clippy::result_large_err)] async fn serve(
829 stream: TcpStream,
830 config: PeerConfig,
831 shared: Arc<Shared>,
832 shutdown: CancellationToken,
833) {
834 let expected = format!("Bearer {}", config.bearer);
835 let inspecting = Arc::clone(&shared);
836 let upgraded = accept_hdr_async(stream, move |request: &Request, response: Response| {
837 inspect_upgrade(request, response, &expected, &inspecting)
838 })
839 .await;
840 let Ok(mut socket) = upgraded else {
841 return;
843 };
844
845 if config.stall == Some(StallPoint::Upgrade) {
846 shutdown.cancelled().await;
849 return;
850 }
851
852 let (directives, mut inbox) = mpsc::channel::<Directive>(16);
853 let _own = directives.clone();
856 let generation = shared.register(directives);
857 let mut session = Session::default();
858
859 if config.withhold != Withhold::SessionCreated {
860 let created = json!({
861 "type": "session.created",
862 "event_id": session.next_event_id(),
863 "session": {"id": "sess_fixture", "object": "realtime.session", "type": "realtime"},
864 });
865 if write(&mut socket, created).await == Next::End {
866 finish(&shared, generation);
867 return;
868 }
869 }
870
871 loop {
872 let step = tokio::select! {
873 biased;
874 () = shutdown.cancelled() => Step::Stop,
875 frame = socket.next() => Step::Inbound(frame),
876 directive = inbox.recv() => Step::Directed(directive),
877 };
878 let next = match step {
879 Step::Stop | Step::Directed(None) => Next::End,
880 Step::Inbound(frame) => {
881 inbound(frame, &mut socket, &config, &shared, &mut session).await
882 }
883 Step::Directed(Some(directive)) => {
884 directed(directive, &mut socket, &config, &shared, &mut session).await
885 }
886 };
887 match next {
888 Next::Serve => {}
889 Next::End => break,
890 Next::Stall => {
891 shutdown.cancelled().await;
896 break;
897 }
898 }
899 }
900 finish(&shared, generation);
901}
902
903enum Step {
905 Stop,
906 Inbound(Option<Result<Message, tokio_tungstenite::tungstenite::Error>>),
907 Directed(Option<Directive>),
908}
909
910#[derive(Debug, Clone, Copy, PartialEq, Eq)]
912enum Next {
913 Serve,
915 End,
917 Stall,
919}
920
921fn finish(shared: &Arc<Shared>, generation: u64) {
922 shared.unregister(generation);
923 shared.update(|record| record.sessions_ended += 1);
924}
925
926#[allow(clippy::result_large_err)]
929fn inspect_upgrade(
930 request: &Request,
931 response: Response,
932 expected: &str,
933 shared: &Arc<Shared>,
934) -> Result<Response, ErrorResponse> {
935 let authorization = request
936 .headers()
937 .get("authorization")
938 .and_then(|value| value.to_str().ok())
939 .map(str::to_owned);
940 let authorised = authorization.as_deref() == Some(expected);
941 let upgrade = Upgrade {
942 target: request.uri().to_string(),
943 authorization,
944 header_names: request
945 .headers()
946 .keys()
947 .map(|name| name.as_str().to_owned())
948 .collect(),
949 outcome: if authorised {
950 UpgradeOutcome::Accepted
951 } else {
952 UpgradeOutcome::Refused(401)
953 },
954 };
955 shared.update(move |record| record.upgrades.push(upgrade));
956 if authorised {
957 Ok(response)
958 } else {
959 let mut refusal = ErrorResponse::new(Some(
963 "invalid_request_error: the bearer token is missing or does not match".to_owned(),
964 ));
965 *refusal.status_mut() = StatusCode::UNAUTHORIZED;
966 Err(refusal)
967 }
968}
969
970async fn inbound(
971 frame: Option<Result<Message, tokio_tungstenite::tungstenite::Error>>,
972 socket: &mut Socket,
973 config: &PeerConfig,
974 shared: &Arc<Shared>,
975 session: &mut Session,
976) -> Next {
977 match frame {
978 Some(Ok(Message::Text(text))) => {
979 let event = read_client_event(&text);
980 match &event {
981 ClientEvent::Cancel => {
982 session.cancelled = session
983 .in_flight
984 .clone()
985 .map_or(Cancelled::Pending, Cancelled::Response);
986 }
987 ClientEvent::Append { audio } => {
988 let audio = audio.clone();
989 shared.update(|record| record.appended_audio.extend_from_slice(&audio));
990 }
991 _ => {}
992 }
993 let reply_wanted = matches!(event, ClientEvent::SessionUpdate(_))
994 && config.withhold != Withhold::SessionUpdated;
995 let stall_now = config.stall == Some(StallPoint::Session)
996 && matches!(event, ClientEvent::Append { .. });
997 shared.update(move |record| record.client_events.push(event));
998 if stall_now {
999 return Next::Stall;
1003 }
1004 if reply_wanted {
1005 let updated = json!({
1006 "type": "session.updated",
1007 "event_id": session.next_event_id(),
1008 "session": {"id": "sess_fixture", "object": "realtime.session", "type": "realtime"},
1009 });
1010 return write(socket, updated).await;
1011 }
1012 Next::Serve
1013 }
1014 Some(Ok(Message::Binary(bytes))) => {
1015 shared.update(move |record| {
1016 record.client_events.push(ClientEvent::Unreadable {
1017 reason: format!("a binary frame of {} bytes", bytes.len()),
1018 });
1019 });
1020 Next::Serve
1021 }
1022 Some(Ok(Message::Ping(_))) => {
1023 shared.update(|record| record.pings += 1);
1024 let _flushed = socket.flush().await;
1027 Next::Serve
1028 }
1029 Some(Ok(Message::Pong(_) | Message::Frame(_))) => Next::Serve,
1030 Some(Ok(Message::Close(_)) | Err(_)) | None => Next::End,
1031 }
1032}
1033
1034fn read_client_event(text: &str) -> ClientEvent {
1035 let Ok(event) = serde_json::from_str::<Value>(text) else {
1036 return ClientEvent::Unreadable {
1037 reason: "a text frame that is not JSON".to_owned(),
1038 };
1039 };
1040 let Some(event_type) = event.get("type").and_then(Value::as_str) else {
1041 return ClientEvent::Unreadable {
1042 reason: "a JSON frame with no string `type`".to_owned(),
1043 };
1044 };
1045 match event_type {
1046 "session.update" => ClientEvent::SessionUpdate(event.clone()),
1047 "response.cancel" => ClientEvent::Cancel,
1048 "input_audio_buffer.append" => match event
1049 .get("audio")
1050 .and_then(Value::as_str)
1051 .map(|audio| BASE64.decode(audio))
1052 {
1053 Some(Ok(audio)) => ClientEvent::Append { audio },
1054 Some(Err(error)) => ClientEvent::Unreadable {
1055 reason: format!("an append whose audio is not RFC 4648 §4 base64: {error}"),
1056 },
1057 None => ClientEvent::Unreadable {
1058 reason: "an append with no string `audio` member".to_owned(),
1059 },
1060 },
1061 other => ClientEvent::Outside {
1062 event_type: other.to_owned(),
1063 },
1064 }
1065}
1066
1067async fn directed(
1068 directive: Directive,
1069 socket: &mut Socket,
1070 config: &PeerConfig,
1071 shared: &Arc<Shared>,
1072 session: &mut Session,
1073) -> Next {
1074 let Directive { action, done } = directive;
1075 match action {
1076 Action::Delta { response, audio } => {
1077 if config.cancel == CancelPolicy::Truncate && session.is_cancelled(&response) {
1078 shared.update(|record| record.deltas_suppressed += 1);
1079 let _answered = done.send(Emission::SuppressedByCancel);
1080 return Next::Serve;
1081 }
1082 session.in_flight = Some(response.clone());
1083 let delta = json!({
1084 "type": "response.output_audio.delta",
1085 "event_id": session.next_event_id(),
1086 "response_id": response,
1087 "item_id": "item_fixture",
1088 "output_index": 0,
1089 "content_index": 0,
1090 "delta": BASE64.encode(&audio),
1091 });
1092 let flow = write(socket, delta).await;
1093 if flow == Next::Serve {
1094 shared.update(|record| record.deltas_sent += 1);
1095 }
1096 answer(flow, done)
1097 }
1098 Action::Scripted(scripted) => {
1099 let event = scripted_event(scripted, session);
1100 answer(write(socket, event).await, done)
1101 }
1102 Action::Malformed(malformed) => malformed_frame(malformed, socket, session, done).await,
1103 Action::Close { code } => {
1104 let close = Message::Close(Some(CloseFrame {
1105 code: CloseCode::from(code),
1106 reason: Utf8Bytes::from_static("session ended"),
1107 }));
1108 let sent = socket.send(close).await.is_ok();
1109 let _flushed = socket.flush().await;
1110 let _answered = done.send(Emission::Sent);
1111 if sent { Next::Serve } else { Next::End }
1114 }
1115 Action::Reset => {
1116 #[allow(deprecated)]
1124 let _lingered = socket.get_ref().set_linger(Some(Duration::ZERO));
1125 let _answered = done.send(Emission::Sent);
1126 Next::End
1127 }
1128 }
1129}
1130
1131fn scripted_event(scripted: Scripted, session: &mut Session) -> Value {
1137 match scripted {
1138 Scripted::AudioDone { response } => {
1139 let mut event = json!({
1140 "type": "response.output_audio.done",
1141 "event_id": session.next_event_id(),
1142 "item_id": "item_fixture",
1143 "output_index": 0,
1144 "content_index": 0,
1145 });
1146 if let (Some(response), Some(members)) = (response, event.as_object_mut()) {
1147 members.insert("response_id".to_owned(), Value::String(response));
1148 }
1149 event
1150 }
1151 Scripted::ResponseDone { response, status } => {
1152 session.in_flight = None;
1153 session.cancelled = Cancelled::No;
1154 json!({
1155 "type": "response.done",
1156 "event_id": session.next_event_id(),
1157 "response": {"id": response, "object": "realtime.response", "status": status},
1158 })
1159 }
1160 Scripted::SpeechStarted => json!({
1161 "type": "input_audio_buffer.speech_started",
1162 "event_id": session.next_event_id(),
1163 "audio_start_ms": 460,
1164 "item_id": "item_fixture",
1165 }),
1166 Scripted::Error { code, message } => json!({
1167 "type": "error",
1168 "event_id": session.next_event_id(),
1169 "error": {
1170 "type": "invalid_request_error",
1171 "code": code,
1172 "message": message,
1173 "param": Value::Null,
1174 },
1175 }),
1176 Scripted::Unknown { event_type } => {
1177 json!({"type": event_type, "event_id": session.next_event_id()})
1178 }
1179 Scripted::Oversize { bytes } => {
1180 json!({
1184 "type": "response.output_audio.delta",
1185 "event_id": session.next_event_id(),
1186 "response_id": "resp_oversize",
1187 "delta": "A".repeat(bytes.next_multiple_of(4)),
1188 })
1189 }
1190 }
1191}
1192
1193async fn malformed_frame(
1194 malformed: Malformed,
1195 socket: &mut Socket,
1196 session: &mut Session,
1197 done: oneshot::Sender<Emission>,
1198) -> Next {
1199 let flow = match malformed {
1200 Malformed::NotJson => socket
1201 .send(Message::Text(Utf8Bytes::from_static("not json{")))
1202 .await
1203 .map_or(Next::End, |()| Next::Serve),
1204 Malformed::NoType => {
1205 let event = json!({"event_id": session.next_event_id(), "session": {}});
1206 write(socket, event).await
1207 }
1208 Malformed::Binary => socket
1209 .send(Message::Binary(WsBytes::from_static(b"\x00\x01binary")))
1210 .await
1211 .map_or(Next::End, |()| Next::Serve),
1212 Malformed::DeltaNotBase64 { response } => {
1213 let event = json!({
1214 "type": "response.output_audio.delta",
1215 "event_id": session.next_event_id(),
1216 "response_id": response,
1217 "delta": "not base64!!",
1218 });
1219 write(socket, event).await
1220 }
1221 Malformed::DeltaMissing { response } => {
1222 let event = json!({
1223 "type": "response.output_audio.delta",
1224 "event_id": session.next_event_id(),
1225 "response_id": response,
1226 "item_id": "item_fixture",
1227 });
1228 write(socket, event).await
1229 }
1230 Malformed::AudioDoneWithoutResponseId => {
1231 let event = json!({
1232 "type": "response.output_audio.done",
1233 "event_id": session.next_event_id(),
1234 "item_id": "item_fixture",
1235 });
1236 write(socket, event).await
1237 }
1238 };
1239 answer(flow, done)
1240}
1241
1242fn answer(flow: Next, done: oneshot::Sender<Emission>) -> Next {
1248 if flow == Next::Serve {
1249 let _answered = done.send(Emission::Sent);
1250 }
1251 flow
1252}
1253
1254async fn write(socket: &mut Socket, event: Value) -> Next {
1257 if socket
1258 .send(Message::Text(event.to_string().into()))
1259 .await
1260 .is_err()
1261 {
1262 return Next::End;
1263 }
1264 match socket.flush().await {
1265 Ok(()) => Next::Serve,
1266 Err(_) => Next::End,
1267 }
1268}
1269
1270impl fmt::Display for Emission {
1271 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
1272 formatter.write_str(match self {
1273 Self::Sent => "sent",
1274 Self::SuppressedByCancel => "suppressed by cancel",
1275 })
1276 }
1277}