use super::client::{
ExampleStoreClient, ExamplesArrayFilter, SearchExampleResult, SearchExamplesRequest,
};
use adk_core::{BeforeModelCallback, BeforeModelResult, Content, Part, Result};
use std::sync::Arc;
const DEFAULT_TOP_K: i64 = 5;
pub struct ExampleStoreProvider {
client: Arc<ExampleStoreClient>,
top_k: i64,
function_names: Option<ExamplesArrayFilter>,
fail_open: bool,
}
impl ExampleStoreProvider {
pub fn new(client: Arc<ExampleStoreClient>) -> Self {
Self { client, top_k: DEFAULT_TOP_K, function_names: None, fail_open: false }
}
#[must_use]
pub fn with_top_k(mut self, top_k: i64) -> Self {
self.top_k = top_k;
self
}
#[must_use]
pub fn with_function_names(mut self, function_names: ExamplesArrayFilter) -> Self {
self.function_names = Some(function_names);
self
}
#[must_use]
pub fn fail_open(mut self, fail_open: bool) -> Self {
self.fail_open = fail_open;
self
}
pub async fn retrieve(&self, query: &str) -> Result<Vec<SearchExampleResult>> {
let mut request = SearchExamplesRequest::by_search_key(query, self.top_k);
if let Some(function_names) = &self.function_names {
request = request.with_function_names(function_names.clone());
}
Ok(self.client.search_examples(request).await?.results)
}
pub fn format_examples(results: &[SearchExampleResult]) -> String {
let mut block = String::from(
"The following retrieved examples show how to respond to similar requests:\n",
);
for (index, result) in results.iter().enumerate() {
block.push_str(&format!("\nExample {}:\n", index + 1));
let contents_example = &result.example.stored_contents_example.contents_example;
for content in &contents_example.contents {
append_content_lines(&mut block, content);
}
for expected in &contents_example.expected_contents {
append_content_lines(&mut block, &expected.content);
}
}
block
}
pub fn into_before_model_callback(self) -> BeforeModelCallback {
let provider = Arc::new(self);
Box::new(move |_ctx, mut request| {
let provider = provider.clone();
Box::pin(async move {
let Some(query) = last_user_text(&request.contents) else {
return Ok(BeforeModelResult::Continue(request));
};
match provider.retrieve(&query).await {
Ok(results) => {
if !results.is_empty() {
let block = Self::format_examples(&results);
request.contents.insert(0, Content::new("user").with_text(block));
}
Ok(BeforeModelResult::Continue(request))
}
Err(error) if provider.fail_open => {
tracing::warn!(
error = %error,
"example store retrieval failed — continuing without examples"
);
Ok(BeforeModelResult::Continue(request))
}
Err(error) => Err(error),
}
})
})
}
}
fn append_content_lines(block: &mut String, content: &Content) {
for part in &content.parts {
if let Part::Text { text } = part {
block.push_str(&format!(" {}: {text}\n", content.role));
}
}
}
fn last_user_text(contents: &[Content]) -> Option<String> {
contents.iter().rev().find(|content| content.role == "user").and_then(|content| {
content
.parts
.iter()
.find_map(|part| if let Part::Text { text } = part { Some(text.clone()) } else { None })
})
}
#[cfg(test)]
mod tests {
use super::super::client::{ContentsExample, Example, StoredContentsExample};
use super::*;
fn result_with_turns(user: &str, model: &str) -> SearchExampleResult {
SearchExampleResult {
example: Example::new(StoredContentsExample::new(ContentsExample::new(
vec![Content::new("user").with_text(user)],
vec![Content::new("model").with_text(model)],
))),
similarity_score: Some(0.9),
}
}
#[test]
fn test_format_examples_renders_role_prefixed_turns() {
let results =
vec![result_with_turns("What is 2+2?", "4"), result_with_turns("Capital?", "Paris")];
let block = ExampleStoreProvider::format_examples(&results);
assert_eq!(
block,
"The following retrieved examples show how to respond to similar requests:\n\
\nExample 1:\n user: What is 2+2?\n model: 4\n\
\nExample 2:\n user: Capital?\n model: Paris\n",
);
}
#[test]
fn test_last_user_text_picks_the_most_recent_user_content() {
let contents = vec![
Content::new("user").with_text("instruction preamble"),
Content::new("model").with_text("hello"),
Content::new("user").with_text("actual question"),
];
assert_eq!(last_user_text(&contents), Some("actual question".to_string()));
assert_eq!(last_user_text(&[]), None);
}
}