gemini_genai_rs/transport/
replay.rs1use std::collections::VecDeque;
19use std::sync::Arc;
20
21use async_trait::async_trait;
22use tokio::sync::watch;
23
24use super::recording::{WireDirection, WireEntry};
25use super::ws::Transport;
26
27pub type OutboundFrames = Arc<parking_lot::Mutex<Vec<Vec<u8>>>>;
29
30pub type FrameGate = Arc<
35 dyn Fn(u64) -> std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send>> + Send + Sync,
36>;
37
38#[derive(Debug, thiserror::Error)]
40pub enum ReplayTransportError {
41 #[error("Not connected")]
43 NotConnected,
44}
45
46#[derive(Clone)]
49pub struct ReplayControl {
50 gate_tx: Arc<watch::Sender<bool>>,
51 drained_rx: watch::Receiver<bool>,
52 outbound: OutboundFrames,
53}
54
55impl ReplayControl {
56 pub fn release(&self) {
60 let _ = self.gate_tx.send(true);
61 }
62
63 pub async fn drained(&self) {
68 let mut rx = self.drained_rx.clone();
69 while !*rx.borrow() {
70 if rx.changed().await.is_err() {
71 break;
72 }
73 }
74 }
75
76 pub fn outbound_frames(&self) -> Vec<Vec<u8>> {
78 self.outbound.lock().clone()
79 }
80}
81
82pub struct ReplayTransport {
85 inbound: VecDeque<Vec<u8>>,
86 timestamps: VecDeque<u64>,
87 on_frame: Option<FrameGate>,
88 ungated_prefix: usize,
89 delivered: usize,
90 gate_rx: watch::Receiver<bool>,
91 drained_tx: watch::Sender<bool>,
92 outbound: OutboundFrames,
93 connected: bool,
94}
95
96impl ReplayTransport {
97 pub fn from_frames(frames: Vec<Vec<u8>>) -> (Self, ReplayControl) {
102 let (gate_tx, gate_rx) = watch::channel(false);
103 let (drained_tx, drained_rx) = watch::channel(false);
104 let outbound: OutboundFrames = Arc::new(parking_lot::Mutex::new(Vec::new()));
105 let control = ReplayControl {
106 gate_tx: Arc::new(gate_tx),
107 drained_rx,
108 outbound: outbound.clone(),
109 };
110 (
111 Self {
112 inbound: frames.into(),
113 timestamps: VecDeque::new(),
114 on_frame: None,
115 ungated_prefix: 1,
116 delivered: 0,
117 gate_rx,
118 drained_tx,
119 outbound,
120 connected: false,
121 },
122 control,
123 )
124 }
125
126 pub fn from_wire_log(entries: &[WireEntry]) -> (Self, ReplayControl) {
132 let inbound: Vec<&WireEntry> = entries
133 .iter()
134 .filter(|e| e.dir == WireDirection::Inbound)
135 .collect();
136 let (mut transport, control) =
137 Self::from_frames(inbound.iter().map(|e| e.payload.clone()).collect());
138 transport.timestamps = inbound.iter().map(|e| e.ts_ms).collect();
139 (transport, control)
140 }
141
142 pub fn with_frame_gate(mut self, gate: FrameGate) -> Self {
147 self.on_frame = Some(gate);
148 self
149 }
150
151 pub fn with_ungated_prefix(mut self, n: usize) -> Self {
154 self.ungated_prefix = n;
155 self
156 }
157}
158
159#[async_trait]
160impl Transport for ReplayTransport {
161 type Error = ReplayTransportError;
162
163 async fn connect(
164 &mut self,
165 _url: &str,
166 _headers: Vec<(String, String)>,
167 ) -> Result<(), Self::Error> {
168 self.connected = true;
169 Ok(())
170 }
171
172 async fn send(&mut self, data: Vec<u8>) -> Result<(), Self::Error> {
173 if !self.connected {
174 return Err(ReplayTransportError::NotConnected);
175 }
176 self.outbound.lock().push(data);
177 Ok(())
178 }
179
180 async fn recv(&mut self) -> Result<Option<Vec<u8>>, Self::Error> {
181 if !self.connected {
182 return Err(ReplayTransportError::NotConnected);
183 }
184 tokio::task::yield_now().await;
187
188 if self.inbound.is_empty() {
189 let _ = self.drained_tx.send(true);
190 std::future::pending::<()>().await;
193 unreachable!("pending() never resolves");
194 }
195
196 if self.delivered >= self.ungated_prefix {
197 let mut gate = self.gate_rx.clone();
198 while !*gate.borrow() {
199 if gate.changed().await.is_err() {
202 break;
203 }
204 }
205 }
206
207 let frame = self
208 .inbound
209 .pop_front()
210 .expect("checked non-empty inbound queue");
211 if let Some(ts_ms) = self.timestamps.pop_front()
212 && let Some(gate) = &self.on_frame
213 {
214 gate(ts_ms).await;
215 }
216 self.delivered += 1;
217 if self.inbound.is_empty() {
218 let _ = self.drained_tx.send(true);
219 }
220 Ok(Some(frame))
221 }
222
223 async fn close(&mut self) -> Result<(), Self::Error> {
224 self.connected = false;
225 Ok(())
226 }
227}
228
229#[cfg(test)]
230mod tests {
231 use super::*;
232 use std::time::Duration;
233
234 #[tokio::test]
235 async fn replay_delivers_first_frame_ungated_then_waits_for_release() {
236 let (mut transport, control) = ReplayTransport::from_frames(vec![
237 br#"{"setupComplete":{}}"#.to_vec(),
238 br#"{"serverContent":{"turnComplete":true}}"#.to_vec(),
239 ]);
240 transport.connect("replay://", vec![]).await.unwrap();
241
242 let first = transport.recv().await.unwrap().unwrap();
244 assert!(String::from_utf8(first).unwrap().contains("setupComplete"));
245
246 let gated = tokio::time::timeout(Duration::from_millis(50), transport.recv()).await;
248 assert!(gated.is_err(), "second frame should be gated");
249
250 control.release();
251 let second = transport.recv().await.unwrap().unwrap();
252 assert!(String::from_utf8(second).unwrap().contains("turnComplete"));
253
254 tokio::time::timeout(Duration::from_millis(100), control.drained())
256 .await
257 .expect("drained should be signalled");
258
259 let idle = tokio::time::timeout(Duration::from_millis(50), transport.recv()).await;
261 assert!(idle.is_err(), "recv should pend after drain");
262 }
263
264 #[tokio::test]
265 async fn replay_collects_outbound_frames() {
266 let (mut transport, control) =
267 ReplayTransport::from_frames(vec![br#"{"setupComplete":{}}"#.to_vec()]);
268 transport.connect("replay://", vec![]).await.unwrap();
269 transport.send(b"{\"setup\":{}}".to_vec()).await.unwrap();
270 transport
271 .send(b"{\"toolResponse\":{}}".to_vec())
272 .await
273 .unwrap();
274
275 let sent = control.outbound_frames();
276 assert_eq!(sent.len(), 2);
277 assert_eq!(sent[0], b"{\"setup\":{}}".to_vec());
278 }
279
280 #[tokio::test]
281 async fn replay_from_wire_log_keeps_inbound_only() {
282 let entries = vec![
283 WireEntry {
284 seq: 1,
285 dir: WireDirection::Outbound,
286 ts_ms: 1,
287 payload: b"{\"setup\":{}}".to_vec(),
288 },
289 WireEntry {
290 seq: 2,
291 dir: WireDirection::Inbound,
292 ts_ms: 2,
293 payload: br#"{"setupComplete":{}}"#.to_vec(),
294 },
295 ];
296 let (mut transport, _control) = ReplayTransport::from_wire_log(&entries);
297 transport.connect("replay://", vec![]).await.unwrap();
298 let first = transport.recv().await.unwrap().unwrap();
299 assert!(String::from_utf8(first).unwrap().contains("setupComplete"));
300 }
301
302 #[tokio::test]
303 async fn replay_errors_when_not_connected() {
304 let (mut transport, _control) = ReplayTransport::from_frames(vec![]);
305 assert!(transport.recv().await.is_err());
306 assert!(transport.send(vec![1]).await.is_err());
307 }
308}