use std::sync::{Arc, Mutex};
use adk_core::{
Artifacts, CallbackContext, Content, EventActions, MemoryEntry, ReadonlyContext, ToolContext,
};
use adk_memory::MemoryService;
use async_trait::async_trait;
pub struct RealtimeToolContext {
app_name: String,
user_id: String,
session_id: String,
function_call_id: String,
memory_service: Option<Arc<dyn MemoryService>>,
actions: Mutex<EventActions>,
user_content: Content,
}
impl RealtimeToolContext {
pub fn new(
app_name: String,
user_id: String,
session_id: String,
function_call_id: String,
memory_service: Option<Arc<dyn MemoryService>>,
) -> Self {
Self {
app_name,
user_id,
session_id,
function_call_id,
memory_service,
actions: Mutex::new(EventActions::default()),
user_content: Content::new("user"),
}
}
}
#[async_trait]
impl ReadonlyContext for RealtimeToolContext {
fn invocation_id(&self) -> &str {
&self.function_call_id
}
fn agent_name(&self) -> &str {
"realtime"
}
fn user_id(&self) -> &str {
&self.user_id
}
fn app_name(&self) -> &str {
&self.app_name
}
fn session_id(&self) -> &str {
&self.session_id
}
fn branch(&self) -> &str {
"main"
}
fn user_content(&self) -> &Content {
&self.user_content
}
}
#[async_trait]
impl CallbackContext for RealtimeToolContext {
fn artifacts(&self) -> Option<Arc<dyn Artifacts>> {
None
}
}
#[async_trait]
impl ToolContext for RealtimeToolContext {
fn function_call_id(&self) -> &str {
&self.function_call_id
}
fn actions(&self) -> EventActions {
self.actions.lock().unwrap().clone()
}
fn set_actions(&self, actions: EventActions) {
*self.actions.lock().unwrap() = actions;
}
async fn search_memory(&self, query: &str) -> adk_core::Result<Vec<MemoryEntry>> {
match &self.memory_service {
Some(service) => {
let resp = service
.search(adk_memory::SearchRequest {
query: query.to_string(),
user_id: self.user_id.clone(),
app_name: self.app_name.clone(),
limit: Some(10),
min_score: None,
project_id: None,
})
.await?;
Ok(resp
.memories
.into_iter()
.map(|m| MemoryEntry { content: m.content, author: m.author })
.collect())
}
None => Ok(vec![]),
}
}
}
use super::SessionIdentity;
pub trait ToolContextFactory: Send + Sync {
fn create_context(&self, function_call_id: &str) -> Arc<dyn ToolContext>;
}
pub struct DefaultToolContextFactory {
pub identity: SessionIdentity,
pub memory_service: Option<Arc<dyn MemoryService>>,
}
impl ToolContextFactory for DefaultToolContextFactory {
fn create_context(&self, function_call_id: &str) -> Arc<dyn ToolContext> {
Arc::new(RealtimeToolContext::new(
self.identity.app_name.clone(),
self.identity.user_id.clone(),
self.identity.session_id.clone(),
function_call_id.to_string(),
self.memory_service.clone(),
))
}
}