1use std::sync::Arc;
31
32use async_trait::async_trait;
33use parking_lot::Mutex;
34
35use gemini_genai_rs::prelude::{Content, FunctionResponse};
36use gemini_genai_rs::session::{SessionError, SessionWriter};
37
38pub struct PendingContext {
49 buffer: Mutex<Vec<Content>>,
50 prompt: Mutex<bool>,
52}
53
54impl PendingContext {
55 pub fn new() -> Self {
57 Self {
58 buffer: Mutex::new(Vec::new()),
59 prompt: Mutex::new(false),
60 }
61 }
62
63 pub fn push(&self, content: Content) {
65 self.buffer.lock().push(content);
66 }
67
68 pub fn extend(&self, contents: Vec<Content>) {
70 if !contents.is_empty() {
71 self.buffer.lock().extend(contents);
72 }
73 }
74
75 pub fn set_prompt(&self) {
77 *self.prompt.lock() = true;
78 }
79
80 pub fn drain(&self) -> (Vec<Content>, bool) {
84 let contents = self.drain_context();
85 let prompt = self.take_prompt();
86 (contents, prompt)
87 }
88
89 pub fn drain_context(&self) -> Vec<Content> {
91 {
92 let mut buf = self.buffer.lock();
93 std::mem::take(&mut *buf)
94 }
95 }
96
97 pub fn take_prompt(&self) -> bool {
99 let mut p = self.prompt.lock();
100 std::mem::replace(&mut *p, false)
101 }
102
103 pub fn clear_prompt(&self) {
105 *self.prompt.lock() = false;
106 }
107
108 pub fn has_prompt(&self) -> bool {
110 *self.prompt.lock()
111 }
112
113 pub fn is_empty(&self) -> bool {
115 self.buffer.lock().is_empty() && !*self.prompt.lock()
116 }
117}
118
119impl Default for PendingContext {
120 fn default() -> Self {
121 Self::new()
122 }
123}
124
125pub struct DeferredWriter {
150 inner: Arc<dyn SessionWriter>,
151 pending: Arc<PendingContext>,
152}
153
154impl DeferredWriter {
155 pub fn new(inner: Arc<dyn SessionWriter>, pending: Arc<PendingContext>) -> Self {
157 Self { inner, pending }
158 }
159
160 async fn flush_context(&self) -> Result<(), SessionError> {
165 let contents = self.pending.drain_context();
166 if !contents.is_empty() {
167 self.inner.send_client_content(contents, false).await?;
168 }
169 Ok(())
170 }
171
172 pub fn pending(&self) -> &Arc<PendingContext> {
174 &self.pending
175 }
176}
177
178#[async_trait]
179impl SessionWriter for DeferredWriter {
180 async fn send_audio(&self, data: bytes::Bytes) -> Result<(), SessionError> {
181 self.flush_context().await?;
182 self.inner.send_audio(data).await
183 }
184
185 async fn send_text(&self, text: String) -> Result<(), SessionError> {
186 self.flush_context().await?;
187 self.inner.send_text(text).await
188 }
189
190 async fn send_tool_response(
191 &self,
192 responses: Vec<FunctionResponse>,
193 ) -> Result<(), SessionError> {
194 self.inner.send_tool_response(responses).await
196 }
197
198 async fn send_client_content(
199 &self,
200 turns: Vec<Content>,
201 turn_complete: bool,
202 ) -> Result<(), SessionError> {
203 self.inner.send_client_content(turns, turn_complete).await
206 }
207
208 async fn send_video(&self, jpeg_data: bytes::Bytes) -> Result<(), SessionError> {
209 self.flush_context().await?;
210 self.inner.send_video(jpeg_data).await
211 }
212
213 async fn update_instruction(&self, instruction: String) -> Result<(), SessionError> {
214 self.inner.update_instruction(instruction).await
216 }
217
218 async fn signal_activity_start(&self) -> Result<(), SessionError> {
219 self.inner.signal_activity_start().await
220 }
221
222 async fn signal_activity_end(&self) -> Result<(), SessionError> {
223 self.inner.signal_activity_end().await
224 }
225
226 async fn disconnect(&self) -> Result<(), SessionError> {
227 let _ = self.flush_context().await;
229 self.inner.disconnect().await
230 }
231}
232
233#[cfg(test)]
234mod tests {
235 use super::*;
236 use std::sync::atomic::{AtomicUsize, Ordering};
237
238 struct CountingWriter {
240 audio_count: AtomicUsize,
241 text_count: AtomicUsize,
242 client_content_count: AtomicUsize,
243 video_count: AtomicUsize,
244 }
245
246 impl CountingWriter {
247 fn new() -> Self {
248 Self {
249 audio_count: AtomicUsize::new(0),
250 text_count: AtomicUsize::new(0),
251 client_content_count: AtomicUsize::new(0),
252 video_count: AtomicUsize::new(0),
253 }
254 }
255 }
256
257 #[async_trait]
258 impl SessionWriter for CountingWriter {
259 async fn send_audio(&self, _: bytes::Bytes) -> Result<(), SessionError> {
260 self.audio_count.fetch_add(1, Ordering::SeqCst);
261 Ok(())
262 }
263 async fn send_text(&self, _: String) -> Result<(), SessionError> {
264 self.text_count.fetch_add(1, Ordering::SeqCst);
265 Ok(())
266 }
267 async fn send_tool_response(&self, _: Vec<FunctionResponse>) -> Result<(), SessionError> {
268 Ok(())
269 }
270 async fn send_client_content(&self, _: Vec<Content>, _: bool) -> Result<(), SessionError> {
271 self.client_content_count.fetch_add(1, Ordering::SeqCst);
272 Ok(())
273 }
274 async fn send_video(&self, _: bytes::Bytes) -> Result<(), SessionError> {
275 self.video_count.fetch_add(1, Ordering::SeqCst);
276 Ok(())
277 }
278 async fn update_instruction(&self, _: String) -> Result<(), SessionError> {
279 Ok(())
280 }
281 async fn signal_activity_start(&self) -> Result<(), SessionError> {
282 Ok(())
283 }
284 async fn signal_activity_end(&self) -> Result<(), SessionError> {
285 Ok(())
286 }
287 async fn disconnect(&self) -> Result<(), SessionError> {
288 Ok(())
289 }
290 }
291
292 #[test]
293 fn pending_context_push_and_drain() {
294 let pc = PendingContext::new();
295 assert!(pc.is_empty());
296
297 pc.push(Content::model("context 1"));
298 pc.push(Content::model("context 2"));
299 assert!(!pc.is_empty());
300
301 let (contents, prompt) = pc.drain();
302 assert_eq!(contents.len(), 2);
303 assert!(!prompt);
304 assert!(pc.is_empty());
305 }
306
307 #[test]
308 fn pending_context_extend() {
309 let pc = PendingContext::new();
310 pc.extend(vec![
311 Content::model("a"),
312 Content::model("b"),
313 Content::model("c"),
314 ]);
315 let (contents, _) = pc.drain();
316 assert_eq!(contents.len(), 3);
317 }
318
319 #[test]
320 fn pending_context_prompt_flag() {
321 let pc = PendingContext::new();
322 pc.push(Content::model("ctx"));
323 pc.set_prompt();
324 assert!(!pc.is_empty());
325
326 let (contents, prompt) = pc.drain();
327 assert_eq!(contents.len(), 1);
328 assert!(prompt);
329 assert!(pc.is_empty());
330 }
331
332 #[test]
333 fn pending_context_drain_clears() {
334 let pc = PendingContext::new();
335 pc.push(Content::model("a"));
336 pc.set_prompt();
337 let _ = pc.drain();
338
339 let (contents, prompt) = pc.drain();
341 assert!(contents.is_empty());
342 assert!(!prompt);
343 }
344
345 #[tokio::test]
346 async fn deferred_writer_flushes_on_send_audio() {
347 let inner = Arc::new(CountingWriter::new());
348 let pending = Arc::new(PendingContext::new());
349 let writer = DeferredWriter::new(inner.clone(), pending.clone());
350
351 pending.push(Content::model("steering context"));
352 pending.push(Content::model("phase instruction"));
353
354 writer.send_audio(vec![0u8; 100].into()).await.unwrap();
355
356 assert_eq!(inner.client_content_count.load(Ordering::SeqCst), 1);
358 assert_eq!(inner.audio_count.load(Ordering::SeqCst), 1);
359 assert!(pending.is_empty());
360 }
361
362 #[tokio::test]
363 async fn deferred_writer_flushes_on_send_text() {
364 let inner = Arc::new(CountingWriter::new());
365 let pending = Arc::new(PendingContext::new());
366 let writer = DeferredWriter::new(inner.clone(), pending.clone());
367
368 pending.push(Content::model("context"));
369
370 writer.send_text("hello".into()).await.unwrap();
371
372 assert_eq!(inner.client_content_count.load(Ordering::SeqCst), 1);
373 assert_eq!(inner.text_count.load(Ordering::SeqCst), 1);
374 }
375
376 #[tokio::test]
377 async fn deferred_writer_flushes_on_send_video() {
378 let inner = Arc::new(CountingWriter::new());
379 let pending = Arc::new(PendingContext::new());
380 let writer = DeferredWriter::new(inner.clone(), pending.clone());
381
382 pending.push(Content::model("context"));
383
384 writer.send_video(vec![0xFFu8; 50].into()).await.unwrap();
385
386 assert_eq!(inner.client_content_count.load(Ordering::SeqCst), 1);
387 assert_eq!(inner.video_count.load(Ordering::SeqCst), 1);
388 }
389
390 #[tokio::test]
391 async fn deferred_writer_no_flush_when_empty() {
392 let inner = Arc::new(CountingWriter::new());
393 let pending = Arc::new(PendingContext::new());
394 let writer = DeferredWriter::new(inner.clone(), pending.clone());
395
396 writer.send_audio(vec![0u8; 100].into()).await.unwrap();
398
399 assert_eq!(inner.client_content_count.load(Ordering::SeqCst), 0);
400 assert_eq!(inner.audio_count.load(Ordering::SeqCst), 1);
401 }
402
403 #[tokio::test]
404 async fn deferred_writer_keeps_prompt_pending_on_user_audio() {
405 let inner = Arc::new(CountingWriter::new());
406 let pending = Arc::new(PendingContext::new());
407 let writer = DeferredWriter::new(inner.clone(), pending.clone());
408
409 pending.push(Content::model("repair nudge"));
410 pending.set_prompt();
411
412 writer.send_audio(vec![0u8; 100].into()).await.unwrap();
413
414 assert_eq!(inner.client_content_count.load(Ordering::SeqCst), 1);
417 assert_eq!(inner.audio_count.load(Ordering::SeqCst), 1);
418 assert!(!pending.is_empty());
419 assert!(pending.take_prompt());
420 }
421
422 #[tokio::test]
423 async fn deferred_writer_does_not_flush_on_tool_response() {
424 let inner = Arc::new(CountingWriter::new());
425 let pending = Arc::new(PendingContext::new());
426 let writer = DeferredWriter::new(inner.clone(), pending.clone());
427
428 pending.push(Content::model("context"));
429
430 writer.send_tool_response(vec![]).await.unwrap();
431
432 assert_eq!(inner.client_content_count.load(Ordering::SeqCst), 0);
434 assert!(!pending.is_empty());
435 }
436
437 #[tokio::test]
438 async fn deferred_writer_client_content_passes_through() {
439 let inner = Arc::new(CountingWriter::new());
440 let pending = Arc::new(PendingContext::new());
441 let writer = DeferredWriter::new(inner.clone(), pending.clone());
442
443 pending.push(Content::model("queued context"));
444
445 writer
447 .send_client_content(vec![Content::user("explicit")], true)
448 .await
449 .unwrap();
450
451 assert_eq!(inner.client_content_count.load(Ordering::SeqCst), 1);
452 assert!(!pending.is_empty());
454 }
455}