1use crate::config::{RealtimeConfig, ToolDefinition, VadConfig, VadMode};
56use crate::events::{ServerEvent, ToolResponse};
57use adk_core::{
58 AdkError, AfterAgentCallback, AfterToolCallback, Agent, BeforeAgentCallback,
59 BeforeToolCallback, CallbackContext, Content, Event, EventActions, EventStream,
60 GlobalInstructionProvider, InstructionProvider, InvocationContext, MemoryEntry, Part,
61 ReadonlyContext, Result, Tool, ToolCallbackContext, ToolContext, Toolset,
62};
63use async_stream::stream;
64use async_trait::async_trait;
65
66use std::sync::{Arc, Mutex};
67
68pub type BoxedRealtimeModel = Arc<dyn crate::model::RealtimeModel>;
70
71pub struct RealtimeAgent {
77 name: String,
78 description: String,
79 model: BoxedRealtimeModel,
80
81 instruction: Option<String>,
83 instruction_provider: Option<Arc<InstructionProvider>>,
84 global_instruction: Option<String>,
85 global_instruction_provider: Option<Arc<GlobalInstructionProvider>>,
86
87 voice: Option<String>,
89 vad_config: Option<VadConfig>,
90 modalities: Vec<String>,
91
92 tools: Vec<Arc<dyn Tool>>,
94 toolsets: Vec<Arc<dyn Toolset>>,
95 sub_agents: Vec<Arc<dyn Agent>>,
96
97 before_callbacks: Arc<Vec<BeforeAgentCallback>>,
99 after_callbacks: Arc<Vec<AfterAgentCallback>>,
100 before_tool_callbacks: Arc<Vec<BeforeToolCallback>>,
101 after_tool_callbacks: Arc<Vec<AfterToolCallback>>,
102
103 on_audio: Option<AudioCallback>,
105 on_transcript: Option<TranscriptCallback>,
106 on_speech_started: Option<SpeechCallback>,
107 on_speech_stopped: Option<SpeechCallback>,
108
109 #[cfg(feature = "video-avatar")]
111 avatar_config: Option<crate::avatar::AvatarConfig>,
112
113 #[cfg(feature = "video-avatar")]
115 avatar_provider: Option<std::sync::Arc<dyn crate::avatar::AvatarProvider>>,
116}
117
118pub type AudioCallback = Arc<
120 dyn Fn(&[u8], &str) -> std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send>>
121 + Send
122 + Sync,
123>;
124
125pub type TranscriptCallback = Arc<
127 dyn Fn(&str, &str) -> std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send>>
128 + Send
129 + Sync,
130>;
131
132pub type SpeechCallback = Arc<
134 dyn Fn(u64) -> std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send>> + Send + Sync,
135>;
136
137impl std::fmt::Debug for RealtimeAgent {
138 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
139 f.debug_struct("RealtimeAgent")
140 .field("name", &self.name)
141 .field("description", &self.description)
142 .field("model", &self.model.model_id())
143 .field("voice", &self.voice)
144 .field("tools_count", &self.tools.len())
145 .field("toolsets_count", &self.toolsets.len())
146 .field("sub_agents_count", &self.sub_agents.len())
147 .finish()
148 }
149}
150
151pub struct RealtimeAgentBuilder {
153 name: String,
154 description: Option<String>,
155 model: Option<BoxedRealtimeModel>,
156 instruction: Option<String>,
157 instruction_provider: Option<Arc<InstructionProvider>>,
158 global_instruction: Option<String>,
159 global_instruction_provider: Option<Arc<GlobalInstructionProvider>>,
160 voice: Option<String>,
161 vad_config: Option<VadConfig>,
162 modalities: Vec<String>,
163 tools: Vec<Arc<dyn Tool>>,
164 toolsets: Vec<Arc<dyn Toolset>>,
165 sub_agents: Vec<Arc<dyn Agent>>,
166 before_callbacks: Vec<BeforeAgentCallback>,
167 after_callbacks: Vec<AfterAgentCallback>,
168 before_tool_callbacks: Vec<BeforeToolCallback>,
169 after_tool_callbacks: Vec<AfterToolCallback>,
170 on_audio: Option<AudioCallback>,
171 on_transcript: Option<TranscriptCallback>,
172 on_speech_started: Option<SpeechCallback>,
173 on_speech_stopped: Option<SpeechCallback>,
174
175 #[cfg(feature = "video-avatar")]
176 avatar_config: Option<crate::avatar::AvatarConfig>,
177
178 #[cfg(feature = "video-avatar")]
179 avatar_provider: Option<std::sync::Arc<dyn crate::avatar::AvatarProvider>>,
180}
181
182impl RealtimeAgentBuilder {
183 pub fn new(name: impl Into<String>) -> Self {
185 Self {
186 name: name.into(),
187 description: None,
188 model: None,
189 instruction: None,
190 instruction_provider: None,
191 global_instruction: None,
192 global_instruction_provider: None,
193 voice: None,
194 vad_config: None,
195 modalities: vec!["text".to_string(), "audio".to_string()],
196 tools: Vec::new(),
197 toolsets: Vec::new(),
198 sub_agents: Vec::new(),
199 before_callbacks: Vec::new(),
200 after_callbacks: Vec::new(),
201 before_tool_callbacks: Vec::new(),
202 after_tool_callbacks: Vec::new(),
203 on_audio: None,
204 on_transcript: None,
205 on_speech_started: None,
206 on_speech_stopped: None,
207 #[cfg(feature = "video-avatar")]
208 avatar_config: None,
209 #[cfg(feature = "video-avatar")]
210 avatar_provider: None,
211 }
212 }
213
214 pub fn description(mut self, desc: impl Into<String>) -> Self {
216 self.description = Some(desc.into());
217 self
218 }
219
220 pub fn model(mut self, model: BoxedRealtimeModel) -> Self {
222 self.model = Some(model);
223 self
224 }
225
226 pub fn instruction(mut self, instruction: impl Into<String>) -> Self {
228 self.instruction = Some(instruction.into());
229 self
230 }
231
232 pub fn instruction_provider(mut self, provider: InstructionProvider) -> Self {
234 self.instruction_provider = Some(Arc::new(provider));
235 self
236 }
237
238 pub fn global_instruction(mut self, instruction: impl Into<String>) -> Self {
240 self.global_instruction = Some(instruction.into());
241 self
242 }
243
244 pub fn global_instruction_provider(mut self, provider: GlobalInstructionProvider) -> Self {
246 self.global_instruction_provider = Some(Arc::new(provider));
247 self
248 }
249
250 pub fn voice(mut self, voice: impl Into<String>) -> Self {
252 self.voice = Some(voice.into());
253 self
254 }
255
256 pub fn vad(mut self, config: VadConfig) -> Self {
258 self.vad_config = Some(config);
259 self
260 }
261
262 pub fn server_vad(mut self) -> Self {
264 self.vad_config = Some(VadConfig {
265 mode: VadMode::ServerVad,
266 threshold: Some(0.5),
267 prefix_padding_ms: Some(300),
268 silence_duration_ms: Some(500),
269 interrupt_response: Some(true),
270 eagerness: None,
271 });
272 self
273 }
274
275 pub fn modalities(mut self, modalities: Vec<String>) -> Self {
277 self.modalities = modalities;
278 self
279 }
280
281 pub fn tool(mut self, tool: Arc<dyn Tool>) -> Self {
283 self.tools.push(tool);
284 self
285 }
286
287 pub fn toolset(mut self, toolset: Arc<dyn Toolset>) -> Self {
293 self.toolsets.push(toolset);
294 self
295 }
296
297 pub fn sub_agent(mut self, agent: Arc<dyn Agent>) -> Self {
299 self.sub_agents.push(agent);
300 self
301 }
302
303 pub fn before_agent_callback(mut self, callback: BeforeAgentCallback) -> Self {
305 self.before_callbacks.push(callback);
306 self
307 }
308
309 pub fn after_agent_callback(mut self, callback: AfterAgentCallback) -> Self {
311 self.after_callbacks.push(callback);
312 self
313 }
314
315 pub fn before_tool_callback(mut self, callback: BeforeToolCallback) -> Self {
317 self.before_tool_callbacks.push(callback);
318 self
319 }
320
321 pub fn after_tool_callback(mut self, callback: AfterToolCallback) -> Self {
323 self.after_tool_callbacks.push(callback);
324 self
325 }
326
327 pub fn on_audio(mut self, callback: AudioCallback) -> Self {
329 self.on_audio = Some(callback);
330 self
331 }
332
333 pub fn on_transcript(mut self, callback: TranscriptCallback) -> Self {
335 self.on_transcript = Some(callback);
336 self
337 }
338
339 pub fn on_speech_started(mut self, callback: SpeechCallback) -> Self {
341 self.on_speech_started = Some(callback);
342 self
343 }
344
345 pub fn on_speech_stopped(mut self, callback: SpeechCallback) -> Self {
347 self.on_speech_stopped = Some(callback);
348 self
349 }
350
351 #[cfg(feature = "video-avatar")]
359 pub fn avatar(mut self, config: crate::avatar::AvatarConfig) -> Self {
360 self.avatar_config = Some(config);
361 self
362 }
363
364 #[cfg(feature = "video-avatar")]
385 pub fn avatar_provider(
386 mut self,
387 provider: std::sync::Arc<dyn crate::avatar::AvatarProvider>,
388 ) -> Self {
389 self.avatar_provider = Some(provider);
390 self
391 }
392
393 pub fn build(self) -> Result<RealtimeAgent> {
395 let model =
396 self.model.ok_or_else(|| AdkError::agent("RealtimeModel is required".to_string()))?;
397
398 Ok(RealtimeAgent {
399 name: self.name,
400 description: self.description.unwrap_or_default(),
401 model,
402 instruction: self.instruction,
403 instruction_provider: self.instruction_provider,
404 global_instruction: self.global_instruction,
405 global_instruction_provider: self.global_instruction_provider,
406 voice: self.voice,
407 vad_config: self.vad_config,
408 modalities: self.modalities,
409 tools: self.tools,
410 toolsets: self.toolsets,
411 sub_agents: self.sub_agents,
412 before_callbacks: Arc::new(self.before_callbacks),
413 after_callbacks: Arc::new(self.after_callbacks),
414 before_tool_callbacks: Arc::new(self.before_tool_callbacks),
415 after_tool_callbacks: Arc::new(self.after_tool_callbacks),
416 on_audio: self.on_audio,
417 on_transcript: self.on_transcript,
418 on_speech_started: self.on_speech_started,
419 on_speech_stopped: self.on_speech_stopped,
420 #[cfg(feature = "video-avatar")]
421 avatar_config: self.avatar_config,
422 #[cfg(feature = "video-avatar")]
423 avatar_provider: self.avatar_provider,
424 })
425 }
426}
427
428impl RealtimeAgent {
429 pub fn builder(name: impl Into<String>) -> RealtimeAgentBuilder {
431 RealtimeAgentBuilder::new(name)
432 }
433
434 pub fn instruction(&self) -> Option<&String> {
436 self.instruction.as_ref()
437 }
438
439 pub fn voice(&self) -> Option<&String> {
441 self.voice.as_ref()
442 }
443
444 pub fn vad_config(&self) -> Option<&VadConfig> {
446 self.vad_config.as_ref()
447 }
448
449 pub fn tools(&self) -> &[Arc<dyn Tool>] {
451 &self.tools
452 }
453
454 #[cfg(feature = "video-avatar")]
458 pub fn avatar_config(&self) -> Option<&crate::avatar::AvatarConfig> {
459 self.avatar_config.as_ref()
460 }
461
462 #[cfg(feature = "video-avatar")]
466 pub fn avatar_provider(&self) -> Option<&std::sync::Arc<dyn crate::avatar::AvatarProvider>> {
467 self.avatar_provider.as_ref()
468 }
469
470 async fn build_config(
472 &self,
473 ctx: &Arc<dyn InvocationContext>,
474 resolved_tools: &[Arc<dyn Tool>],
475 ) -> Result<RealtimeConfig> {
476 let mut config = RealtimeConfig::default();
477
478 if let Some(provider) = &self.global_instruction_provider {
480 let global_inst = provider(ctx.clone() as Arc<dyn ReadonlyContext>).await?;
481 if !global_inst.is_empty() {
482 config.instruction = Some(global_inst);
483 }
484 } else if let Some(ref template) = self.global_instruction {
485 let processed = adk_core::inject_session_state(ctx.as_ref(), template).await?;
486 config.instruction = Some(processed);
487 }
488
489 if let Some(provider) = &self.instruction_provider {
491 let inst = provider(ctx.clone() as Arc<dyn ReadonlyContext>).await?;
492 if !inst.is_empty() {
493 if let Some(existing) = &mut config.instruction {
494 existing.push_str("\n\n");
495 existing.push_str(&inst);
496 } else {
497 config.instruction = Some(inst);
498 }
499 }
500 } else if let Some(ref template) = self.instruction {
501 let processed = adk_core::inject_session_state(ctx.as_ref(), template).await?;
502 if let Some(existing) = &mut config.instruction {
503 existing.push_str("\n\n");
504 existing.push_str(&processed);
505 } else {
506 config.instruction = Some(processed);
507 }
508 }
509
510 config.voice = self.voice.clone();
512 config.turn_detection = self.vad_config.clone();
513 config.modalities = Some(self.modalities.clone());
514
515 let tool_defs: Vec<ToolDefinition> = resolved_tools
517 .iter()
518 .map(|t| ToolDefinition {
519 name: t.name().to_string(),
520 description: Some(t.enhanced_description().to_string()),
521 parameters: t.parameters_schema(),
522 })
523 .collect();
524
525 if !tool_defs.is_empty() {
526 config.tools = Some(tool_defs);
527 }
528
529 if !self.sub_agents.is_empty() {
531 let mut tools = config.tools.unwrap_or_default();
532 tools.push(ToolDefinition {
533 name: "transfer_to_agent".to_string(),
534 description: Some("Transfer execution to another agent.".to_string()),
535 parameters: Some(serde_json::json!({
536 "type": "object",
537 "properties": {
538 "agent_name": {
539 "type": "string",
540 "description": "The name of the agent to transfer to."
541 }
542 },
543 "required": ["agent_name"]
544 })),
545 });
546 config.tools = Some(tools);
547 }
548
549 #[cfg(feature = "video-avatar")]
554 if let Some(ref avatar) = self.avatar_config {
555 tracing::warn!(
556 agent = %self.name,
557 source_url = %avatar.source_url,
558 "video avatar configured but the current realtime provider does not support video avatars; proceeding audio-only"
559 );
560 let avatar_json = serde_json::to_value(avatar).unwrap_or_else(|e| {
561 tracing::warn!("failed to serialize avatar config: {e}");
562 serde_json::Value::Null
563 });
564 let extra = config.extra.get_or_insert_with(|| serde_json::json!({}));
565 if let Some(obj) = extra.as_object_mut() {
566 obj.insert("avatarConfig".to_string(), avatar_json);
567 }
568 }
569
570 Ok(config)
571 }
572
573 #[allow(dead_code)]
575 async fn execute_tool(
576 &self,
577 ctx: &Arc<dyn InvocationContext>,
578 call_id: &str,
579 name: &str,
580 arguments: &str,
581 ) -> (serde_json::Value, EventActions) {
582 let tool = self.tools.iter().find(|t| t.name() == name);
584
585 if let Some(tool) = tool {
586 let args: serde_json::Value =
587 serde_json::from_str(arguments).unwrap_or(serde_json::json!({}));
588
589 let tool_ctx: Arc<dyn ToolContext> =
591 Arc::new(RealtimeToolContext::new(ctx.clone(), call_id.to_string()));
592
593 let tool_cb_ctx =
595 Arc::new(ToolCallbackContext::new(ctx.clone(), name.to_string(), args.clone()));
596 for callback in self.before_tool_callbacks.as_ref() {
597 if let Err(e) = callback(tool_cb_ctx.clone() as Arc<dyn CallbackContext>).await {
598 return (
599 serde_json::json!({ "error": e.to_string() }),
600 EventActions::default(),
601 );
602 }
603 }
604
605 let result = match tool.execute(tool_ctx.clone(), args.clone()).await {
607 Ok(result) => result,
608 Err(e) => serde_json::json!({ "error": e.to_string() }),
609 };
610
611 let actions = tool_ctx.actions();
612
613 let tool_cb_ctx =
615 Arc::new(ToolCallbackContext::new(ctx.clone(), name.to_string(), args.clone()));
616 for callback in self.after_tool_callbacks.as_ref() {
617 if let Err(e) = callback(tool_cb_ctx.clone() as Arc<dyn CallbackContext>).await {
618 return (serde_json::json!({ "error": e.to_string() }), actions);
619 }
620 }
621
622 (result, actions)
623 } else {
624 (
625 serde_json::json!({ "error": format!("Tool {} not found", name) }),
626 EventActions::default(),
627 )
628 }
629 }
630}
631
632#[async_trait]
633impl Agent for RealtimeAgent {
634 fn name(&self) -> &str {
635 &self.name
636 }
637
638 fn description(&self) -> &str {
639 &self.description
640 }
641
642 fn sub_agents(&self) -> &[Arc<dyn Agent>] {
643 &self.sub_agents
644 }
645
646 async fn run(&self, ctx: Arc<dyn InvocationContext>) -> Result<EventStream> {
647 let agent_name = self.name.clone();
648 let invocation_id = ctx.invocation_id().to_string();
649 let model = self.model.clone();
650 let _sub_agents = self.sub_agents.clone();
651
652 let before_callbacks = self.before_callbacks.clone();
654 let after_callbacks = self.after_callbacks.clone();
655 let before_tool_callbacks = self.before_tool_callbacks.clone();
656 let after_tool_callbacks = self.after_tool_callbacks.clone();
657 let tools = self.tools.clone();
658 let toolsets = self.toolsets.clone();
659
660 let on_audio = self.on_audio.clone();
662 let on_transcript = self.on_transcript.clone();
663 let on_speech_started = self.on_speech_started.clone();
664 let on_speech_stopped = self.on_speech_stopped.clone();
665
666 #[cfg(feature = "video-avatar")]
668 let avatar_provider = self.avatar_provider.clone();
669 #[cfg(feature = "video-avatar")]
670 let avatar_config_for_session = self.avatar_config.clone();
671
672 let mut resolved_tools: Vec<Arc<dyn Tool>> = tools.clone();
674 let static_tool_names: std::collections::HashSet<String> =
675 tools.iter().map(|t| t.name().to_string()).collect();
676 let mut toolset_source: std::collections::HashMap<String, String> =
677 std::collections::HashMap::new();
678
679 for toolset in &toolsets {
680 let toolset_tools = toolset.tools(ctx.clone() as Arc<dyn ReadonlyContext>).await?;
681 for tool in &toolset_tools {
682 let name = tool.name().to_string();
683 if static_tool_names.contains(&name) {
684 return Err(AdkError::agent(format!(
685 "Duplicate tool name '{}': conflict between static tool and toolset '{}'",
686 name,
687 toolset.name()
688 )));
689 }
690 if let Some(other_toolset_name) = toolset_source.get(&name) {
691 return Err(AdkError::agent(format!(
692 "Duplicate tool name '{}': conflict between toolset '{}' and toolset '{}'",
693 name,
694 other_toolset_name,
695 toolset.name()
696 )));
697 }
698 toolset_source.insert(name, toolset.name().to_string());
699 resolved_tools.push(tool.clone());
700 }
701 }
702
703 let config = self.build_config(&ctx, &resolved_tools).await?;
705
706 let s = stream! {
707 for callback in before_callbacks.as_ref() {
709 match callback(ctx.clone() as Arc<dyn CallbackContext>).await {
710 Ok(Some(content)) => {
711 let mut early_event = Event::new(&invocation_id);
712 early_event.author = agent_name.clone();
713 early_event.llm_response.content = Some(content);
714 yield Ok(early_event);
715 return;
716 }
717 Ok(None) => continue,
718 Err(e) => {
719 yield Err(e);
720 return;
721 }
722 }
723 }
724
725 let session = match model.connect(config).await {
727 Ok(s) => s,
728 Err(e) => {
729 yield Err(AdkError::model(format!("Failed to connect: {}", e)));
730 return;
731 }
732 };
733
734 let mut start_event = Event::new(&invocation_id);
736 start_event.author = agent_name.clone();
737 start_event.llm_response.content = Some(Content {
738 role: "system".to_string(),
739 parts: vec![Part::Text {
740 text: format!("Realtime session started: {}", session.session_id()),
741 }],
742 });
743 yield Ok(start_event);
744
745 #[cfg(feature = "video-avatar")]
747 let avatar_session_id: Option<String> = {
748 if let (Some(provider), Some(config)) = (&avatar_provider, &avatar_config_for_session) {
749 match provider.start_session(config).await {
750 Ok(session_info) => {
751 tracing::info!(
752 provider = %session_info.provider,
753 session_id = %session_info.session_id,
754 "avatar session started"
755 );
756 let mut avatar_event = Event::new(&invocation_id);
758 avatar_event.author = agent_name.clone();
759 avatar_event.llm_response.content = Some(Content {
760 role: "system".to_string(),
761 parts: vec![Part::Text {
762 text: serde_json::to_string(&session_info).unwrap_or_default(),
763 }],
764 });
765 yield Ok(avatar_event);
766 Some(session_info.session_id)
767 }
768 Err(e) => {
769 tracing::warn!(
771 error = %e,
772 "avatar session creation failed, falling back to audio-only"
773 );
774 None
775 }
776 }
777 } else {
778 None
779 }
780 };
781 #[cfg(not(feature = "video-avatar"))]
782 let _avatar_session_id: Option<String> = None;
783
784 #[cfg(feature = "video-avatar")]
786 let _avatar_keep_alive_handle: Option<tokio::task::JoinHandle<()>> = {
787 if let (Some(provider), Some(sess_id)) = (&avatar_provider, &avatar_session_id) {
788 Some(crate::avatar::spawn_keep_alive(
789 provider.clone(),
790 sess_id.clone(),
791 std::time::Duration::from_secs(30),
792 ))
793 } else {
794 None
795 }
796 };
797
798 let user_content = ctx.user_content();
801 for part in &user_content.parts {
802 if let Part::Text { text } = part {
803 if let Err(e) = session.send_text(text).await {
804 yield Err(AdkError::model(format!("Failed to send text: {}", e)));
805 return;
806 }
807 if let Err(e) = session.create_response().await {
809 yield Err(AdkError::model(format!("Failed to create response: {}", e)));
810 return;
811 }
812 }
813 }
814
815 loop {
817 let event = session.next_event().await;
818
819 match event {
820 Some(Ok(server_event)) => {
821 match server_event {
822 ServerEvent::AudioDelta { delta, item_id, .. } => {
823 #[cfg(feature = "video-avatar")]
825 if let (Some(provider), Some(sess_id)) = (&avatar_provider, &avatar_session_id) {
826 if let Err(e) = provider.send_audio(sess_id, &delta).await {
827 tracing::warn!(error = %e, "avatar send_audio failed");
828 }
829 if let Some(ref cb) = on_audio {
832 cb(&delta, &item_id).await;
833 }
834 continue;
835 }
836
837 if let Some(ref cb) = on_audio {
839 cb(&delta, &item_id).await;
840 }
841
842 let mut audio_event = Event::new(&invocation_id);
844 audio_event.author = agent_name.clone();
845 audio_event.llm_response.content = Some(Content {
846 role: "model".to_string(),
847 parts: vec![Part::InlineData {
848 mime_type: "audio/pcm".to_string(),
849 data: delta,
850 uri: None,
851 annotations: None,
852 }],
853 });
854 yield Ok(audio_event);
855 }
856
857 ServerEvent::TextDelta { delta, .. } => {
858 let mut text_event = Event::new(&invocation_id);
859 text_event.author = agent_name.clone();
860 text_event.llm_response.content = Some(Content {
861 role: "model".to_string(),
862 parts: vec![Part::Text { text: delta.clone() }],
863 });
864 yield Ok(text_event);
865 }
866
867 ServerEvent::TranscriptDelta { delta, item_id, .. } => {
868 if let Some(ref cb) = on_transcript {
869 cb(&delta, &item_id).await;
870 }
871 }
872
873 ServerEvent::SpeechStarted { audio_start_ms, .. } => {
874 if let Some(ref cb) = on_speech_started {
875 cb(audio_start_ms).await;
876 }
877 }
878
879 ServerEvent::SpeechStopped { audio_end_ms, .. } => {
880 if let Some(ref cb) = on_speech_stopped {
881 cb(audio_end_ms).await;
882 }
883 }
884
885 ServerEvent::FunctionCallDone {
886 call_id,
887 name,
888 arguments,
889 ..
890 } => {
891 if name == "transfer_to_agent" {
893 let args: serde_json::Value = serde_json::from_str(&arguments)
894 .unwrap_or(serde_json::json!({}));
895 let target = args.get("agent_name")
896 .and_then(|v| v.as_str())
897 .unwrap_or_default()
898 .to_string();
899
900 let mut transfer_event = Event::new(&invocation_id);
901 transfer_event.author = agent_name.clone();
902 transfer_event.actions.transfer_to_agent = Some(target);
903 yield Ok(transfer_event);
904
905 let _ = session.close().await;
906 return;
907 }
908
909 let tool = resolved_tools.iter().find(|t| t.name() == name);
911
912 let (result, actions) = if let Some(tool) = tool {
913 let args: serde_json::Value = serde_json::from_str(&arguments)
914 .unwrap_or(serde_json::json!({}));
915
916 let tool_ctx: Arc<dyn ToolContext> = Arc::new(
917 RealtimeToolContext::new(ctx.clone(), call_id.clone())
918 );
919
920 let cb_ctx: Arc<dyn CallbackContext> =
921 Arc::new(ToolCallbackContext::new(
922 ctx.clone(),
923 name.clone(),
924 args.clone(),
925 ));
926
927 let result = execute_tool_with_callbacks(
928 tool.as_ref(),
929 tool_ctx.clone(),
930 cb_ctx,
931 args.clone(),
932 before_tool_callbacks.as_ref(),
933 after_tool_callbacks.as_ref(),
934 )
935 .await;
936
937 (result, tool_ctx.actions())
938 } else {
939 (
940 serde_json::json!({ "error": format!("Tool {} not found", name) }),
941 EventActions::default(),
942 )
943 };
944
945 let mut tool_event = Event::new(&invocation_id);
947 tool_event.author = agent_name.clone();
948 tool_event.actions = actions.clone();
949 tool_event.llm_response.content = Some(Content {
950 role: "function".to_string(),
951 parts: vec![Part::FunctionResponse {
952 function_response: adk_core::FunctionResponseData::new(name.clone(), result.clone()),
953 id: Some(call_id.clone()),
954 annotations: None,
955 }],
956 });
957 yield Ok(tool_event);
958
959 if actions.escalate || actions.skip_summarization {
961 let _ = session.close().await;
962 return;
963 }
964
965 let response = ToolResponse {
967 call_id,
968 output: result,
969 };
970 if let Err(e) = session.send_tool_response(response).await {
971 yield Err(AdkError::model(format!("Failed to send tool response: {}", e)));
972 let _ = session.close().await;
973 return;
974 }
975 }
976
977 ServerEvent::ResponseDone { .. } => {
978 }
980
981 ServerEvent::Error { error, .. } => {
982 yield Err(AdkError::model(format!(
983 "Realtime error: {} - {}",
984 error.code.unwrap_or_default(),
985 error.message
986 )));
987 }
988
989
990 _ => {
991 }
993 }
994 }
995 Some(Err(e)) => {
996 yield Err(AdkError::model(format!("Session error: {}", e)));
997 break;
998 }
999 None => {
1000 break;
1002 }
1003 }
1004 }
1005
1006 #[cfg(feature = "video-avatar")]
1008 {
1009 if let Some(handle) = _avatar_keep_alive_handle {
1011 handle.abort();
1012 }
1013 if let (Some(provider), Some(sess_id)) = (&avatar_provider, &avatar_session_id) {
1015 if let Err(e) = provider.stop_session(sess_id).await {
1016 tracing::warn!(error = %e, "avatar session cleanup failed");
1017 }
1018 }
1019 }
1020
1021 for callback in after_callbacks.as_ref() {
1023 match callback(ctx.clone() as Arc<dyn CallbackContext>).await {
1024 Ok(Some(content)) => {
1025 let mut after_event = Event::new(&invocation_id);
1026 after_event.author = agent_name.clone();
1027 after_event.llm_response.content = Some(content);
1028 yield Ok(after_event);
1029 break;
1030 }
1031 Ok(None) => continue,
1032 Err(e) => {
1033 yield Err(e);
1034 return;
1035 }
1036 }
1037 }
1038 };
1039
1040 Ok(Box::pin(s))
1041 }
1042}
1043
1044async fn execute_tool_with_callbacks(
1055 tool: &dyn Tool,
1056 tool_ctx: Arc<dyn ToolContext>,
1057 cb_ctx: Arc<dyn CallbackContext>,
1058 args: serde_json::Value,
1059 before_tool_callbacks: &[BeforeToolCallback],
1060 after_tool_callbacks: &[AfterToolCallback],
1061) -> serde_json::Value {
1062 let mut short_circuit: Option<serde_json::Value> = None;
1063 let mut run_after_tool_callbacks = true;
1064
1065 for callback in before_tool_callbacks {
1066 match callback(cb_ctx.clone()).await {
1067 Ok(Some(content)) => {
1068 short_circuit = Some(content_to_tool_result(&content));
1069 break;
1070 }
1071 Ok(None) => continue,
1072 Err(e) => {
1073 short_circuit = Some(serde_json::json!({ "error": e.to_string() }));
1074 run_after_tool_callbacks = false;
1075 break;
1076 }
1077 }
1078 }
1079
1080 let mut result = match short_circuit {
1081 Some(result) => result,
1082 None => match tool.execute(tool_ctx, args).await {
1083 Ok(value) => value,
1084 Err(e) => serde_json::json!({ "error": e.to_string() }),
1085 },
1086 };
1087
1088 if run_after_tool_callbacks {
1089 for callback in after_tool_callbacks {
1090 match callback(cb_ctx.clone()).await {
1091 Ok(Some(modified)) => {
1092 result = content_to_tool_result(&modified);
1093 break;
1094 }
1095 Ok(None) => continue,
1096 Err(e) => {
1097 result = serde_json::json!({ "error": e.to_string() });
1098 break;
1099 }
1100 }
1101 }
1102 }
1103
1104 result
1105}
1106
1107fn content_to_tool_result(content: &Content) -> serde_json::Value {
1113 for part in &content.parts {
1114 if let Part::FunctionResponse { function_response, .. } = part {
1115 return function_response.response.clone();
1116 }
1117 }
1118
1119 let text: String = content.parts.iter().filter_map(|part| part.text()).collect();
1120 serde_json::json!({ "result": text })
1121}
1122
1123struct RealtimeToolContext {
1124 parent_ctx: Arc<dyn InvocationContext>,
1125 function_call_id: String,
1126 actions: Mutex<EventActions>,
1127}
1128
1129impl RealtimeToolContext {
1130 fn new(parent_ctx: Arc<dyn InvocationContext>, function_call_id: String) -> Self {
1131 Self { parent_ctx, function_call_id, actions: Mutex::new(EventActions::default()) }
1132 }
1133}
1134
1135#[async_trait]
1136impl ReadonlyContext for RealtimeToolContext {
1137 fn invocation_id(&self) -> &str {
1138 self.parent_ctx.invocation_id()
1139 }
1140
1141 fn agent_name(&self) -> &str {
1142 self.parent_ctx.agent_name()
1143 }
1144
1145 fn user_id(&self) -> &str {
1146 self.parent_ctx.user_id()
1147 }
1148
1149 fn app_name(&self) -> &str {
1150 self.parent_ctx.app_name()
1151 }
1152
1153 fn session_id(&self) -> &str {
1154 self.parent_ctx.session_id()
1155 }
1156
1157 fn branch(&self) -> &str {
1158 self.parent_ctx.branch()
1159 }
1160
1161 fn user_content(&self) -> &Content {
1162 self.parent_ctx.user_content()
1163 }
1164}
1165
1166#[async_trait]
1167impl CallbackContext for RealtimeToolContext {
1168 fn artifacts(&self) -> Option<Arc<dyn adk_core::Artifacts>> {
1169 self.parent_ctx.artifacts()
1170 }
1171
1172 fn shared_state(&self) -> Option<Arc<adk_core::SharedState>> {
1175 self.parent_ctx.shared_state()
1176 }
1177}
1178
1179#[async_trait]
1180impl ToolContext for RealtimeToolContext {
1181 fn function_call_id(&self) -> &str {
1182 &self.function_call_id
1183 }
1184
1185 fn actions(&self) -> EventActions {
1186 self.actions.lock().unwrap().clone()
1187 }
1188
1189 fn set_actions(&self, actions: EventActions) {
1190 *self.actions.lock().unwrap() = actions;
1191 }
1192
1193 async fn search_memory(&self, query: &str) -> Result<Vec<MemoryEntry>> {
1194 if let Some(memory) = self.parent_ctx.memory() {
1195 memory.search(query).await
1196 } else {
1197 Ok(vec![])
1198 }
1199 }
1200
1201 fn user_scopes(&self) -> Vec<String> {
1206 self.parent_ctx.user_scopes()
1207 }
1208
1209 async fn get_secret(&self, name: &str) -> Result<Option<String>> {
1214 self.parent_ctx.get_secret(name).await
1215 }
1216}
1217
1218#[cfg(test)]
1219mod tool_safety_tests {
1220 use super::*;
1231 use adk_core::{RunConfig, SharedState, State};
1232 use std::collections::HashMap;
1233 use std::sync::atomic::{AtomicUsize, Ordering};
1234
1235 struct CountingTool {
1237 executions: Arc<AtomicUsize>,
1238 }
1239
1240 #[async_trait]
1241 impl Tool for CountingTool {
1242 fn name(&self) -> &str {
1243 "counting"
1244 }
1245 fn description(&self) -> &str {
1246 "counts executions"
1247 }
1248 async fn execute(
1249 &self,
1250 _ctx: Arc<dyn ToolContext>,
1251 _args: serde_json::Value,
1252 ) -> Result<serde_json::Value> {
1253 self.executions.fetch_add(1, Ordering::SeqCst);
1254 Ok(serde_json::json!({ "ran": true }))
1255 }
1256 }
1257
1258 struct TestToolContext {
1260 actions: Mutex<EventActions>,
1261 content: Content,
1262 }
1263
1264 impl TestToolContext {
1265 fn new() -> Self {
1266 Self { actions: Mutex::new(EventActions::default()), content: Content::new("user") }
1267 }
1268 }
1269
1270 #[async_trait]
1271 impl ReadonlyContext for TestToolContext {
1272 fn invocation_id(&self) -> &str {
1273 "inv"
1274 }
1275 fn agent_name(&self) -> &str {
1276 "agent"
1277 }
1278 fn user_id(&self) -> &str {
1279 "user"
1280 }
1281 fn app_name(&self) -> &str {
1282 "app"
1283 }
1284 fn session_id(&self) -> &str {
1285 "session"
1286 }
1287 fn branch(&self) -> &str {
1288 ""
1289 }
1290 fn user_content(&self) -> &Content {
1291 &self.content
1292 }
1293 }
1294
1295 #[async_trait]
1296 impl CallbackContext for TestToolContext {
1297 fn artifacts(&self) -> Option<Arc<dyn adk_core::Artifacts>> {
1298 None
1299 }
1300 }
1301
1302 #[async_trait]
1303 impl ToolContext for TestToolContext {
1304 fn function_call_id(&self) -> &str {
1305 "call-1"
1306 }
1307 fn actions(&self) -> EventActions {
1308 self.actions.lock().unwrap().clone()
1309 }
1310 fn set_actions(&self, actions: EventActions) {
1311 *self.actions.lock().unwrap() = actions;
1312 }
1313 async fn search_memory(&self, _query: &str) -> Result<Vec<MemoryEntry>> {
1314 Ok(vec![])
1315 }
1316 }
1317
1318 async fn dispatch(
1320 before: Vec<BeforeToolCallback>,
1321 after: Vec<AfterToolCallback>,
1322 executions: Arc<AtomicUsize>,
1323 ) -> serde_json::Value {
1324 let tool = CountingTool { executions };
1325 let ctx = Arc::new(TestToolContext::new());
1326 execute_tool_with_callbacks(
1327 &tool,
1328 ctx.clone() as Arc<dyn ToolContext>,
1329 ctx as Arc<dyn CallbackContext>,
1330 serde_json::json!({}),
1331 &before,
1332 &after,
1333 )
1334 .await
1335 }
1336
1337 #[tokio::test]
1338 async fn a_before_callback_error_prevents_execution() {
1339 let executions = Arc::new(AtomicUsize::new(0));
1340 let before: Vec<BeforeToolCallback> =
1341 vec![Box::new(|_ctx| Box::pin(async { Err(AdkError::tool("denied by policy")) }))];
1342
1343 let result = dispatch(before, vec![], Arc::clone(&executions)).await;
1344
1345 assert_eq!(executions.load(Ordering::SeqCst), 0, "a refused tool must not run: {result}");
1346 assert!(
1347 result["error"].as_str().unwrap_or_default().contains("denied by policy"),
1348 "the refusal reason must reach the provider: {result}"
1349 );
1350 }
1351
1352 #[tokio::test]
1353 async fn a_before_callback_substitution_prevents_execution() {
1354 let executions = Arc::new(AtomicUsize::new(0));
1355 let before: Vec<BeforeToolCallback> = vec![Box::new(|_ctx| {
1356 Box::pin(async {
1357 Ok(Some(Content {
1358 role: "function".to_string(),
1359 parts: vec![Part::FunctionResponse {
1360 function_response: adk_core::FunctionResponseData::new(
1361 "counting",
1362 serde_json::json!({ "cached": true }),
1363 ),
1364 id: None,
1365 annotations: None,
1366 }],
1367 }))
1368 })
1369 })];
1370
1371 let result = dispatch(before, vec![], Arc::clone(&executions)).await;
1372
1373 assert_eq!(
1374 executions.load(Ordering::SeqCst),
1375 0,
1376 "a substituted result must not run the tool"
1377 );
1378 assert_eq!(result, serde_json::json!({ "cached": true }));
1379 }
1380
1381 #[tokio::test]
1382 async fn a_permitting_callback_lets_the_tool_run() {
1383 let executions = Arc::new(AtomicUsize::new(0));
1384 let before: Vec<BeforeToolCallback> = vec![Box::new(|_ctx| Box::pin(async { Ok(None) }))];
1385
1386 let result = dispatch(before, vec![], Arc::clone(&executions)).await;
1387
1388 assert_eq!(executions.load(Ordering::SeqCst), 1);
1389 assert_eq!(result, serde_json::json!({ "ran": true }));
1390 }
1391
1392 #[tokio::test]
1393 async fn an_after_callback_error_becomes_the_result() {
1394 let executions = Arc::new(AtomicUsize::new(0));
1395 let after: Vec<AfterToolCallback> =
1396 vec![Box::new(|_ctx| Box::pin(async { Err(AdkError::tool("post-check failed")) }))];
1397
1398 let result = dispatch(vec![], after, Arc::clone(&executions)).await;
1399
1400 assert_eq!(executions.load(Ordering::SeqCst), 1, "the tool ran, as it should have");
1401 assert!(
1402 result["error"].as_str().unwrap_or_default().contains("post-check failed"),
1403 "an after-callback failure must not be dropped: {result}"
1404 );
1405 }
1406
1407 #[tokio::test]
1408 async fn after_callbacks_are_skipped_when_a_before_callback_refuses() {
1409 let executions = Arc::new(AtomicUsize::new(0));
1410 let after_ran = Arc::new(AtomicUsize::new(0));
1411 let counter = Arc::clone(&after_ran);
1412
1413 let before: Vec<BeforeToolCallback> =
1414 vec![Box::new(|_ctx| Box::pin(async { Err(AdkError::tool("refused")) }))];
1415 let after: Vec<AfterToolCallback> = vec![Box::new(move |_ctx| {
1416 let counter = Arc::clone(&counter);
1417 Box::pin(async move {
1418 counter.fetch_add(1, Ordering::SeqCst);
1419 Ok(None)
1420 })
1421 })];
1422
1423 let result = dispatch(before, after, Arc::clone(&executions)).await;
1424
1425 assert_eq!(executions.load(Ordering::SeqCst), 0);
1426 assert_eq!(after_ran.load(Ordering::SeqCst), 0, "matching the standard loop's ordering");
1427 assert!(result["error"].as_str().unwrap_or_default().contains("refused"), "{result}");
1428 }
1429
1430 struct TestState;
1433 impl State for TestState {
1434 fn get(&self, _key: &str) -> Option<serde_json::Value> {
1435 None
1436 }
1437 fn set(&mut self, _key: String, _value: serde_json::Value) {}
1438 fn all(&self) -> HashMap<String, serde_json::Value> {
1439 HashMap::new()
1440 }
1441 }
1442
1443 struct TestSession;
1444 impl adk_core::Session for TestSession {
1445 fn id(&self) -> &str {
1446 "session"
1447 }
1448 fn app_name(&self) -> &str {
1449 "app"
1450 }
1451 fn user_id(&self) -> &str {
1452 "user"
1453 }
1454 fn state(&self) -> &dyn State {
1455 &TestState
1456 }
1457 fn conversation_history(&self) -> Vec<Content> {
1458 Vec::new()
1459 }
1460 }
1461
1462 struct CapableParent {
1464 content: Content,
1465 config: RunConfig,
1466 session: TestSession,
1467 shared: Arc<SharedState>,
1468 }
1469
1470 #[async_trait]
1471 impl ReadonlyContext for CapableParent {
1472 fn invocation_id(&self) -> &str {
1473 "inv"
1474 }
1475 fn agent_name(&self) -> &str {
1476 "agent"
1477 }
1478 fn user_id(&self) -> &str {
1479 "user"
1480 }
1481 fn app_name(&self) -> &str {
1482 "app"
1483 }
1484 fn session_id(&self) -> &str {
1485 "session"
1486 }
1487 fn branch(&self) -> &str {
1488 ""
1489 }
1490 fn user_content(&self) -> &Content {
1491 &self.content
1492 }
1493 }
1494
1495 #[async_trait]
1496 impl CallbackContext for CapableParent {
1497 fn artifacts(&self) -> Option<Arc<dyn adk_core::Artifacts>> {
1498 None
1499 }
1500
1501 fn shared_state(&self) -> Option<Arc<SharedState>> {
1502 Some(Arc::clone(&self.shared))
1503 }
1504 }
1505
1506 #[async_trait]
1507 impl InvocationContext for CapableParent {
1508 fn agent(&self) -> Arc<dyn Agent> {
1509 unreachable!("not used by these tests")
1510 }
1511 fn memory(&self) -> Option<Arc<dyn adk_core::Memory>> {
1512 None
1513 }
1514 fn session(&self) -> &dyn adk_core::Session {
1515 &self.session
1516 }
1517 fn run_config(&self) -> &RunConfig {
1518 &self.config
1519 }
1520 fn end_invocation(&self) {}
1521 fn ended(&self) -> bool {
1522 false
1523 }
1524 fn user_scopes(&self) -> Vec<String> {
1525 vec!["repo:write".to_string()]
1526 }
1527 async fn get_secret(&self, name: &str) -> Result<Option<String>> {
1528 Ok(Some(format!("secret-for-{name}")))
1529 }
1530 }
1531
1532 #[tokio::test]
1533 async fn the_realtime_tool_context_preserves_parent_capabilities() {
1534 let parent = Arc::new(CapableParent {
1535 content: Content::new("user"),
1536 config: RunConfig::default(),
1537 session: TestSession,
1538 shared: Arc::new(SharedState::new()),
1539 }) as Arc<dyn InvocationContext>;
1540
1541 let ctx = RealtimeToolContext::new(parent, "call-1".to_string());
1542
1543 assert_eq!(
1544 ctx.user_scopes(),
1545 vec!["repo:write".to_string()],
1546 "an empty scope list makes an authenticated caller look anonymous"
1547 );
1548 assert_eq!(ctx.get_secret("api_key").await.unwrap().as_deref(), Some("secret-for-api_key"));
1549 assert!(ctx.shared_state().is_some(), "shared state must reach realtime tools");
1550 assert_eq!(ctx.app_name(), "app", "identity still delegates");
1551 }
1552}