use anda_core::{
AgentOutput, BoxError, BoxPinFut, CompletionRequest, FunctionDefinition, Json, Message,
};
use std::collections::BTreeMap;
use super::driver::{SamplingOptions, WireFormat, drive_completion};
use super::{CompletionFeaturesDyn, ModelEffort, ModelError, request_client_builder};
pub mod types;
impl From<ModelEffort> for types::ThinkingLevel {
fn from(value: ModelEffort) -> Self {
match value {
ModelEffort::Minimal => Self::Minimal,
ModelEffort::Low => Self::Low,
ModelEffort::Medium => Self::Medium,
ModelEffort::High => Self::High,
ModelEffort::Max => Self::High,
}
}
}
fn apply_effort(config: &mut types::GenerationConfig, model: &str, effort: ModelEffort) {
let model = model.trim_start_matches("models/");
let thinking = config.thinking_config.get_or_insert_default();
if model.starts_with("gemini-2.5-") {
let pro = model.contains("-pro");
thinking.thinking_level = None;
thinking.thinking_budget = Some(match effort {
ModelEffort::Minimal => {
if pro {
128
} else {
0
}
}
ModelEffort::Low => 1024,
ModelEffort::Medium => 4096,
ModelEffort::High => 16384,
ModelEffort::Max => {
if pro {
32768
} else {
24576
}
}
});
} else {
thinking.thinking_budget = None;
thinking.thinking_level = Some(match effort {
ModelEffort::Minimal if model.contains("-pro") => types::ThinkingLevel::Low,
ModelEffort::Medium if model.starts_with("gemini-3-pro") => types::ThinkingLevel::High,
effort => effort.into(),
});
}
}
const API_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta/models";
pub static DEFAULT_COMPLETION_MODEL: &str = "gemini-flash-latest";
#[derive(Clone)]
pub struct Client {
endpoint: String,
api_key: String,
http: reqwest::Client,
}
impl Client {
pub fn new(api_key: &str, endpoint: Option<String>) -> Self {
Self::new_with_client(
api_key,
endpoint,
request_client_builder()
.build()
.expect("Gemini reqwest client should build"),
)
}
pub fn new_with_client(api_key: &str, endpoint: Option<String>, http: reqwest::Client) -> Self {
Self {
endpoint: super::resolve_endpoint(endpoint, API_BASE_URL),
api_key: api_key.to_string(),
http,
}
}
pub fn with_client(self, http: reqwest::Client) -> Self {
Self { http, ..self }
}
fn post(&self, path: &str) -> reqwest::RequestBuilder {
let url = format!("{}{}", self.endpoint, path);
self.http.post(url).header("x-goog-api-key", &self.api_key)
}
pub fn completion_model(&self, model: &str) -> CompletionModel {
CompletionModel::new(
self.clone(),
if model.is_empty() {
DEFAULT_COMPLETION_MODEL
} else {
model
},
)
}
}
#[derive(Clone)]
pub struct CompletionModel {
client: Client,
default_request: types::GenerateContentRequest,
pub model: String,
}
impl CompletionModel {
pub fn new(client: Client, model: &str) -> Self {
let default_request = types::GenerateContentRequest::default();
Self {
client,
default_request,
model: model.to_string(),
}
}
pub fn with_stream(mut self, stream: bool) -> Self {
self.default_request.stream = stream;
self
}
pub fn with_effort(mut self, effort: Option<ModelEffort>) -> Self {
if let Some(effort) = effort {
apply_effort(
&mut self.default_request.generation_config,
&self.model,
effort,
);
}
self
}
pub fn with_max_output(mut self, max_output: usize) -> Self {
if max_output > 0 {
self.default_request.generation_config.max_output_tokens =
Some(i32::try_from(max_output).unwrap_or(i32::MAX));
}
self
}
pub fn with_default_request(mut self, greq: types::GenerateContentRequest) -> Self {
self.default_request = greq;
self
}
}
fn append_gemini_parts(target: &mut Vec<types::Part>, parts: Vec<types::Part>) {
for part in parts {
match (target.last_mut(), part) {
(
Some(types::Part {
thought: last_thought,
thought_signature: last_signature,
data: types::PartKind::Text(last_text),
}),
types::Part {
thought,
thought_signature,
data: types::PartKind::Text(text),
},
) if *last_thought == thought
&& last_signature.is_none()
&& thought_signature.is_none() =>
{
last_text.push_str(&text);
}
(_, part) => target.push(part),
}
}
}
fn response_from_stream_chunks(
chunks: Vec<types::GenerateContentResponse>,
) -> Result<types::GenerateContentResponse, BoxError> {
let mut candidates = BTreeMap::<u32, types::Candidate>::new();
let mut prompt_feedback = None;
let mut usage_metadata = types::UsageMetadata::default();
let mut model_version = None;
let mut response_id = None;
let mut model_status = None;
for chunk in chunks {
if chunk.prompt_feedback.is_some() {
prompt_feedback = chunk.prompt_feedback;
}
if chunk.usage_metadata != types::UsageMetadata::default() {
usage_metadata = chunk.usage_metadata;
}
if chunk.model_version.is_some() {
model_version = chunk.model_version;
}
if chunk.response_id.is_some() {
response_id = chunk.response_id;
}
if chunk.model_status.is_some() {
model_status = chunk.model_status;
}
for candidate in chunk.candidates {
let index = candidate.index.unwrap_or(0);
match candidates.entry(index) {
std::collections::btree_map::Entry::Vacant(entry) => {
entry.insert(candidate);
}
std::collections::btree_map::Entry::Occupied(mut entry) => {
let existing = entry.get_mut();
if existing.content.role.is_none() {
existing.content.role = candidate.content.role;
}
append_gemini_parts(&mut existing.content.parts, candidate.content.parts);
if candidate.finish_reason.is_some() {
existing.finish_reason = candidate.finish_reason;
}
if candidate.safety_ratings.is_some() {
existing.safety_ratings = candidate.safety_ratings;
}
if candidate.citation_metadata.is_some() {
existing.citation_metadata = candidate.citation_metadata;
}
if candidate.token_count.is_some() {
existing.token_count = candidate.token_count;
}
if !candidate.grounding_attributions.is_empty() {
existing.grounding_attributions = candidate.grounding_attributions;
}
if candidate.grounding_metadata.is_some() {
existing.grounding_metadata = candidate.grounding_metadata;
}
if candidate.avg_logprobs.is_some() {
existing.avg_logprobs = candidate.avg_logprobs;
}
if candidate.logprobs_result.is_some() {
existing.logprobs_result = candidate.logprobs_result;
}
if candidate.url_context_metadata.is_some() {
existing.url_context_metadata = candidate.url_context_metadata;
}
if candidate.finish_message.is_some() {
existing.finish_message = candidate.finish_message;
}
}
}
}
}
if candidates.is_empty() && prompt_feedback.is_none() {
return Err("No streamed Gemini response".into());
}
if !prompt_feedback
.as_ref()
.is_some_and(types::PromptFeedback::is_blocked)
&& candidates.values().any(|candidate| {
candidate
.finish_reason
.as_ref()
.is_none_or(|reason| matches!(reason, types::FinishReason::FinishReasonUnspecified))
})
{
return Err(Box::new(
ModelError::new("Gemini stream ended before finishReason").with_retryable(true),
));
}
Ok(types::GenerateContentResponse {
candidates: candidates.into_values().collect(),
prompt_feedback,
usage_metadata,
model_version,
response_id,
model_status,
})
}
fn prepare_tool_presentations(message: &mut Message, model: &str) {
if model
.rsplit('/')
.next()
.is_some_and(|model| model.starts_with("gemini-3"))
{
return;
}
for part in &mut message.content {
if let anda_core::ContentPart::ToolOutput { output, .. } = part
&& let Some(view) = anda_core::ToolPresentation::from_output(output)
{
*output = Json::String(view.text_fallback());
}
}
}
impl CompletionFeaturesDyn for CompletionModel {
fn model_name(&self) -> String {
self.model.clone()
}
fn completion(&self, req: CompletionRequest) -> BoxPinFut<Result<AgentOutput, BoxError>> {
let model = self.model.clone();
let client = self.client.clone();
let r = self.default_request.clone();
Box::pin(async move {
drive_completion::<CompletionModel>(model, move |path| client.post(path), r, req).await
})
}
}
impl WireFormat for CompletionModel {
type Request = types::GenerateContentRequest;
type Response = types::GenerateContentResponse;
type StreamItem = types::GenerateContentResponse;
fn set_instructions(r: &mut Self::Request, instructions: String) {
r.system_instruction = Some(types::Content {
role: Some(types::Role::Model),
parts: vec![types::Part {
data: types::PartKind::Text(instructions),
..Default::default()
}],
});
}
fn append_raw_history(r: &mut Self::Request, mut raw_history: Vec<Json>) -> usize {
r.contents.append(&mut raw_history);
r.contents.len()
}
fn push_message(r: &mut Self::Request, mut msg: Message, model: &str) -> Result<(), BoxError> {
prepare_tool_presentations(&mut msg, model);
let val = types::Content::from(msg);
r.contents.push(serde_json::to_value(val)?);
Ok(())
}
fn apply_sampling(
r: &mut Self::Request,
options: SamplingOptions,
model: &str,
) -> Result<(), BoxError> {
if let Some(temperature) = options.temperature {
r.generation_config.temperature = Some(temperature);
}
if let Some(max_tokens) = options.max_output_tokens {
r.generation_config.max_output_tokens = Some(max_tokens as i32);
}
if let Some(effort) = options.effort {
apply_effort(&mut r.generation_config, model, effort);
}
if let Some(output_schema) = options.output_schema {
r.generation_config.response_mime_type = Some("application/json".to_string());
r.generation_config.response_schema = None;
r.generation_config.response_json_schema_compat = None;
r.generation_config.response_json_schema = Some(output_schema);
}
if let Some(stop) = options.stop {
r.generation_config.stop_sequences = Some(stop);
}
Ok(())
}
fn apply_tools(r: &mut Self::Request, tools: Vec<FunctionDefinition>, required: bool) {
r.tools = vec![tools.into()];
let mode = if required {
types::FunctionCallingMode::Any
} else {
types::FunctionCallingMode::Auto
};
r.tool_config = Some(types::ToolConfig {
function_calling_config: types::FunctionCallingConfig {
mode,
allowed_function_names: None,
},
});
}
fn is_stream(r: &Self::Request) -> bool {
r.stream
}
fn endpoint(r: &Self::Request, model: &str) -> String {
if r.stream {
format!("/{}:streamGenerateContent?alt=sse", model)
} else {
format!("/{}:generateContent", model)
}
}
fn aggregate_stream(
items: Vec<Self::StreamItem>,
_done: bool,
) -> Result<Self::Response, BoxError> {
response_from_stream_chunks(items)
}
fn parse_response(
model: &str,
data: &[u8],
) -> Result<(Self::Response, Option<Json>), BoxError> {
match serde_json::from_slice::<types::GenerateContentResponse>(data) {
Ok(res) => Ok((res, None)),
Err(err) => Err(format!(
"Invalid completion response, model: {}, error: {}, body: {}",
model,
err,
super::error_body_excerpt(data)
)
.into()),
}
}
fn maybe_failed(res: &Self::Response) -> bool {
res.maybe_failed()
}
fn sent_messages(mut r: Self::Request, skip_raw: usize) -> Vec<Json> {
if skip_raw > 0 {
r.contents.drain(0..skip_raw);
}
r.contents
}
fn into_output(
res: Self::Response,
sent_messages: Vec<Json>,
chat_history: Vec<Message>,
_assistant_raw_message: Option<Json>,
) -> Result<AgentOutput, BoxError> {
res.try_into(sent_messages, chat_history)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::test_support::{no_proxy_client, recorded, spawn_mock_server, sse_headers};
use anda_core::FunctionDefinition;
use http::{HeaderMap, Method, StatusCode};
use reqwest::header::ACCEPT;
use serde_json::{Value, json};
#[test]
fn thinking_configuration_matches_model_generation() {
for (model, minimal, maximum) in [
("gemini-2.5-pro", 128, 32768),
("gemini-2.5-flash", 0, 24576),
("gemini-2.5-flash-lite", 0, 24576),
] {
let mut request = Client::new("fake", None)
.completion_model(model)
.with_effort(Some(ModelEffort::Minimal))
.default_request;
let config = request.generation_config.thinking_config.as_ref().unwrap();
assert_eq!(config.thinking_budget, Some(minimal));
assert!(config.thinking_level.is_none());
<CompletionModel as WireFormat>::apply_sampling(
&mut request,
SamplingOptions {
effort: Some(ModelEffort::Max),
temperature: None,
max_output_tokens: None,
output_schema: None,
stop: None,
},
model,
)
.unwrap();
assert_eq!(
request
.generation_config
.thinking_config
.unwrap()
.thinking_budget,
Some(maximum)
);
}
let mut config = types::GenerationConfig {
thinking_config: Some(
serde_json::from_value(json!({"includeThoughts": true, "thinkingBudget": -1}))
.unwrap(),
),
..Default::default()
};
assert_eq!(
serde_json::to_value(&config).unwrap()["thinkingConfig"]["thinkingBudget"],
-1
);
apply_effort(&mut config, "gemini-3-flash-preview", ModelEffort::Medium);
let thinking = config.thinking_config.as_ref().unwrap();
assert!(thinking.thinking_budget.is_none());
assert!(thinking.include_thoughts);
assert_eq!(thinking.thinking_level, Some(types::ThinkingLevel::Medium));
apply_effort(&mut config, "gemini-3-pro-preview", ModelEffort::Minimal);
assert_eq!(
config.thinking_config.unwrap().thinking_level,
Some(types::ThinkingLevel::Low)
);
}
#[test]
fn streaming_preserves_signed_parts_and_rejects_unfinished_candidates() {
let parts = vec![
json!({"text":"hello"}),
json!({"text":" world", "thoughtSignature":"sig1"}),
json!({"text":"!", "thoughtSignature":"sig2"}),
json!({"text":"", "thoughtSignature":"sig3"}),
];
let mut chunks: Vec<types::GenerateContentResponse> = parts
.iter()
.map(|part| {
serde_json::from_value(json!({
"candidates":[{"content":{"role":"model","parts":[part]}}]
}))
.unwrap()
})
.collect();
let error = response_from_stream_chunks(chunks.clone()).unwrap_err();
assert!(crate::model::is_retryable_box_error(&error));
chunks
.push(serde_json::from_value(json!({"candidates":[{"finishReason":"STOP"}]})).unwrap());
let output = response_from_stream_chunks(chunks)
.unwrap()
.try_into(vec![], vec![])
.unwrap();
assert_eq!(output.raw_history[0]["parts"], json!(parts));
assert_eq!(output.content, "hello world!");
assert_eq!(
output.chat_history[0].text().as_deref(),
Some("hello world!")
);
let incomplete_call = serde_json::from_value(json!({"candidates":[{"content":{"parts":[{"functionCall":{"name":"lookup","args":{}}}]}}]})).unwrap();
assert!(response_from_stream_chunks(vec![incomplete_call]).is_err());
}
#[test]
fn prompt_safety_feedback_only_fails_when_blocked() {
for feedback in [
json!({}),
json!({"safetyRatings":[]}),
json!({"blockReason":"BLOCK_REASON_UNSPECIFIED"}),
] {
let response: types::GenerateContentResponse = serde_json::from_value(json!({
"promptFeedback":feedback,
"candidates":[{"content":{"role":"model","parts":[{"text":"ok"}]},"finishReason":"STOP"}]
})).unwrap();
assert!(!response.maybe_failed());
let out = response.try_into(vec![], vec![]).unwrap();
assert!(out.failed_reason.is_none());
assert_eq!(out.content, "ok");
}
let blocked =
serde_json::from_value(json!({"promptFeedback":{"blockReason":"SAFETY"}})).unwrap();
let out = response_from_stream_chunks(vec![blocked])
.unwrap()
.try_into(vec![], vec![])
.unwrap();
assert!(out.failed_reason.is_some());
}
async fn complete(
model: &CompletionModel,
req: CompletionRequest,
) -> Result<AgentOutput, BoxError> {
CompletionFeaturesDyn::completion(model, req).await
}
#[test]
fn completion_model_applies_default_effort() {
let model = Client::new("test-key", Some("http://localhost".into()))
.completion_model("gemini-3-pro")
.with_effort(Some(ModelEffort::High));
let thinking_config = model
.default_request
.generation_config
.thinking_config
.expect("thinking config should be configured");
assert_eq!(
thinking_config.thinking_level,
Some(types::ThinkingLevel::High)
);
}
#[test]
fn configured_max_output_replaces_the_default_output_budget() {
let model = Client::new("test-key", None).completion_model("gemini-2.0-flash");
let default_budget = model.default_request.generation_config.max_output_tokens;
assert_eq!(default_budget, Some(65535));
assert_eq!(
model
.clone()
.with_max_output(0)
.default_request
.generation_config
.max_output_tokens,
default_budget
);
assert_eq!(
model
.with_max_output(8192)
.default_request
.generation_config
.max_output_tokens,
Some(8192)
);
}
#[test]
fn client_defaults_effort_mapping_and_default_request_are_covered() {
assert_eq!(
types::ThinkingLevel::from(ModelEffort::Minimal),
types::ThinkingLevel::Minimal
);
assert_eq!(
types::ThinkingLevel::from(ModelEffort::Low),
types::ThinkingLevel::Low
);
assert_eq!(
types::ThinkingLevel::from(ModelEffort::Medium),
types::ThinkingLevel::Medium
);
assert_eq!(
types::ThinkingLevel::from(ModelEffort::High),
types::ThinkingLevel::High
);
assert_eq!(
types::ThinkingLevel::from(ModelEffort::Max),
types::ThinkingLevel::High
);
let default_client = Client::new("test-key", None);
assert_eq!(default_client.endpoint, API_BASE_URL);
let empty_endpoint_client = Client::new("test-key", Some(String::new()));
assert_eq!(empty_endpoint_client.endpoint, API_BASE_URL);
let default_model = empty_endpoint_client.completion_model("");
assert_eq!(default_model.model, DEFAULT_COMPLETION_MODEL);
assert_eq!(
CompletionFeaturesDyn::model_name(&default_model),
DEFAULT_COMPLETION_MODEL
);
let no_effort = default_model.clone().with_effort(None);
assert!(
no_effort
.default_request
.generation_config
.thinking_config
.is_none()
);
let custom_request = types::GenerateContentRequest {
stream: true,
generation_config: types::GenerationConfig {
temperature: Some(0.25),
..Default::default()
},
..Default::default()
};
let custom_model = default_model
.with_stream(false)
.with_default_request(custom_request);
assert!(custom_model.default_request.stream);
assert_eq!(
custom_model.default_request.generation_config.temperature,
Some(0.25)
);
}
#[test]
fn stream_chunk_aggregation_covers_metadata_replacements_and_errors() {
assert_eq!(
response_from_stream_chunks(Vec::new())
.unwrap_err()
.to_string(),
"No streamed Gemini response"
);
let chunks = vec![
serde_json::from_value::<types::GenerateContentResponse>(json!({
"candidates": [{
"index": 0,
"content": {
"parts": [
{"text": "plain"},
{"thought": true, "text": "think"}
]
}
}]
}))
.unwrap(),
serde_json::from_value::<types::GenerateContentResponse>(json!({
"promptFeedback": {"blockReason": "OTHER"},
"modelStatus": {"modelStage": "PREVIEW"},
"candidates": [{
"index": 0,
"content": {
"role": "model",
"parts": [
{
"thought": true,
"thoughtSignature": "sig-1",
"text": " more"
},
{"functionCall": {"name": "lookup", "args": {"q": "anda"}}}
]
},
"finishReason": "MAX_TOKENS",
"citationMetadata": {
"citationSources": [{"uri": "https://example.com"}]
},
"tokenCount": 9,
"groundingAttributions": [{
"sourceId": {"groundingPassage": {"passageId": "p1"}}
}],
"groundingMetadata": {
"webSearchQueries": ["anda coverage"]
},
"avgLogprobs": -0.5,
"logprobsResult": {"logProbabilitySum": -1.0},
"urlContextMetadata": {
"urlMetadata": [{
"retrievedUrl": "https://example.com",
"urlRetrievalStatus": "URL_RETRIEVAL_STATUS_SUCCESS"
}]
},
"finishMessage": "cut short"
}]
}))
.unwrap(),
serde_json::from_value::<types::GenerateContentResponse>(json!({
"candidates": [{
"index": 1,
"content": {
"role": "model",
"parts": [{"text": "second candidate"}]
},
"finishReason": "STOP"
}]
}))
.unwrap(),
];
let response = response_from_stream_chunks(chunks).unwrap();
assert_eq!(
response.prompt_feedback.unwrap().block_reason,
Some(types::BlockReason::Other)
);
assert_eq!(
response.model_status.unwrap().model_stage,
Some(types::ModelStage::Preview)
);
assert_eq!(response.candidates.len(), 2);
let first = &response.candidates[0];
assert_eq!(first.content.role, Some(types::Role::Model));
assert!(matches!(
&first.content.parts[2],
types::Part {
thought: Some(true),
thought_signature: Some(signature),
data: types::PartKind::Text(text),
} if signature == "sig-1" && text == " more"
));
assert!(matches!(
&first.content.parts[3].data,
types::PartKind::FunctionCall { name, args, .. }
if name == "lookup" && args.as_ref() == Some(&json!({"q": "anda"}))
));
assert_eq!(first.finish_reason, Some(types::FinishReason::MaxTokens));
assert_eq!(first.token_count, Some(9));
assert_eq!(first.finish_message.as_deref(), Some("cut short"));
assert_eq!(first.avg_logprobs, Some(-0.5));
assert!(first.citation_metadata.is_some());
assert_eq!(first.grounding_attributions.len(), 1);
assert!(first.grounding_metadata.is_some());
assert!(first.logprobs_result.is_some());
assert!(first.url_context_metadata.is_some());
}
#[tokio::test]
async fn completion_model_posts_request_and_parses_non_stream_response() {
let body = serde_json::to_vec(&json!({
"candidates": [{
"index": 0,
"content": {
"role": "model",
"parts": [{"text": "hello from gemini"}]
},
"finishReason": "STOP"
}],
"usageMetadata": {
"promptTokenCount": 7,
"cachedContentTokenCount": 2,
"candidatesTokenCount": 3
},
"modelVersion": "gemini-test",
"responseId": "resp_1"
}))
.unwrap();
let (endpoint, state) = spawn_mock_server(StatusCode::OK, HeaderMap::new(), body).await;
let model = Client::new("test-key", Some(endpoint))
.with_client(no_proxy_client())
.completion_model("gemini-test")
.with_stream(false);
let output = complete(
&model,
CompletionRequest {
instructions: "system rules".into(),
prompt: "say hello".into(),
temperature: Some(0.4),
max_output_tokens: Some(128),
output_schema: Some(json!({"type": "object"})),
stop: Some(vec!["END".into()]),
effort: Some(ModelEffort::Low),
tools: vec![FunctionDefinition {
name: "lookup".into(),
description: "Lookup docs".into(),
parameters: json!({"type": "object"}),
strict: Some(false),
}],
..Default::default()
},
)
.await
.unwrap();
assert_eq!(output.content, "hello from gemini");
assert_eq!(output.model.as_deref(), Some("gemini-test"));
assert_eq!(output.usage.input_tokens, 7);
assert_eq!(output.usage.cached_tokens, 2);
assert_eq!(output.usage.output_tokens, 3);
let req = recorded(&state);
assert_eq!(req.method, Method::POST);
assert_eq!(req.uri.path(), "/gemini-test:generateContent");
assert_eq!(req.headers.get("x-goog-api-key").unwrap(), "test-key");
assert_ne!(
req.headers.get(ACCEPT).and_then(|v| v.to_str().ok()),
Some("text/event-stream")
);
let sent: Value = serde_json::from_slice(&req.body).unwrap();
assert_eq!(sent["systemInstruction"]["role"], "model");
assert_eq!(sent["contents"][0]["role"], "user");
assert_eq!(sent["generationConfig"]["temperature"], 0.4);
assert_eq!(sent["generationConfig"]["maxOutputTokens"], 128);
assert_eq!(
sent["generationConfig"]["responseMimeType"],
"application/json"
);
assert_eq!(
sent["generationConfig"]["responseJsonSchema"],
json!({"type": "object"})
);
assert_eq!(sent["generationConfig"]["stopSequences"], json!(["END"]));
assert_eq!(
sent["generationConfig"]["thinkingConfig"]["thinkingLevel"],
"LOW"
);
assert_eq!(
sent["tools"][0]["functionDeclarations"][0]["name"],
"lookup"
);
assert!(sent.get("toolConfig").is_some());
}
#[tokio::test]
async fn completion_model_reports_http_and_invalid_json_errors() {
let (endpoint, _) = spawn_mock_server(
StatusCode::SERVICE_UNAVAILABLE,
HeaderMap::new(),
"unavailable",
)
.await;
let model = Client::new("test-key", Some(endpoint))
.with_client(no_proxy_client())
.completion_model("gemini-test")
.with_stream(false);
let err = complete(
&model,
CompletionRequest {
prompt: "hello".into(),
..Default::default()
},
)
.await
.unwrap_err();
assert!(err.to_string().contains("Completion failed"));
assert!(err.to_string().contains("unavailable"));
let (endpoint, _) = spawn_mock_server(StatusCode::OK, HeaderMap::new(), "not json").await;
let model = Client::new("test-key", Some(endpoint))
.with_client(no_proxy_client())
.completion_model("gemini-test")
.with_stream(false);
let err = complete(
&model,
CompletionRequest {
prompt: "hello".into(),
..Default::default()
},
)
.await
.unwrap_err();
assert!(err.to_string().contains("Invalid completion response"));
assert!(err.to_string().contains("not json"));
}
#[tokio::test]
async fn completion_model_streams_sse_chunks() {
let chunks = [
json!({
"candidates": [{
"index": 0,
"content": {
"role": "model",
"parts": [{"text": "Hel"}]
}
}],
"usageMetadata": {"promptTokenCount": 4},
"modelVersion": "gemini-stream"
}),
json!({
"candidates": [{
"index": 0,
"content": {"parts": [{"text": "lo"}]},
"finishReason": "STOP"
}],
"usageMetadata": {
"promptTokenCount": 4,
"candidatesTokenCount": 2
},
"responseId": "resp_stream"
}),
]
.into_iter()
.map(|chunk| format!("data: {chunk}\n\n"))
.collect::<String>();
let (endpoint, state) =
spawn_mock_server(StatusCode::OK, sse_headers(), chunks.into_bytes()).await;
let model = Client::new("test-key", Some(endpoint))
.with_client(no_proxy_client())
.completion_model("gemini-stream")
.with_stream(true);
let output = complete(
&model,
CompletionRequest {
prompt: "stream".into(),
..Default::default()
},
)
.await
.unwrap();
assert_eq!(output.content, "Hello");
assert_eq!(output.usage.input_tokens, 4);
assert_eq!(output.usage.output_tokens, 2);
let req = recorded(&state);
assert_eq!(req.uri.path(), "/gemini-stream:streamGenerateContent");
assert_eq!(req.uri.query(), Some("alt=sse"));
assert_eq!(
req.headers.get(ACCEPT).and_then(|v| v.to_str().ok()),
Some("text/event-stream")
);
assert_eq!(
req.headers
.get(http::header::ACCEPT_ENCODING)
.and_then(|v| v.to_str().ok()),
Some("identity")
);
}
#[test]
fn aggregates_gemini_stream_chunks() {
let chunks = vec![
serde_json::from_value::<types::GenerateContentResponse>(json!({
"candidates": [{
"index": 0,
"content": {
"role": "model",
"parts": [{"text": "Hel"}]
}
}],
"usageMetadata": {"promptTokenCount": 2},
"modelVersion": "gemini-flash-latest"
}))
.unwrap(),
serde_json::from_value::<types::GenerateContentResponse>(json!({
"candidates": [{
"index": 0,
"content": {"parts": [{"text": "lo"}]},
"safetyRatings": [{
"category": "HARM_CATEGORY_VENDOR_SPECIFIC",
"probability": "VERY_LOW"
}]
}]
}))
.unwrap(),
serde_json::from_value::<types::GenerateContentResponse>(json!({
"candidates": [{
"index": 0,
"finishReason": "STOP"
}],
"usageMetadata": {
"promptTokenCount": 2,
"candidatesTokenCount": 1,
"promptTokensDetails": null,
"candidatesTokensDetails": [{
"modality": "FUTURE_MODALITY",
"tokenCount": null
}]
},
"responseId": "resp_1"
}))
.unwrap(),
];
let response = response_from_stream_chunks(chunks).unwrap();
assert!(!response.maybe_failed());
assert_eq!(response.response_id.as_deref(), Some("resp_1"));
assert!(matches!(
&response.candidates[0]
.safety_ratings
.as_ref()
.unwrap()[0]
.category,
types::HarmCategory::Unknown(category)
if category == "HARM_CATEGORY_VENDOR_SPECIFIC"
));
assert!(matches!(
&response.usage_metadata.candidates_tokens_details[0].modality,
Some(types::Modality::Unknown(modality)) if modality == "FUTURE_MODALITY"
));
let output = response.try_into(vec![], vec![]).unwrap();
assert_eq!(output.content, "Hello");
assert_eq!(output.usage.input_tokens, 2);
assert_eq!(output.usage.output_tokens, 1);
assert_eq!(output.chat_history.len(), 1);
}
#[test]
fn prune_inline_media_covers_gemini_inline_parts() {
let bytes = b"inline attachment bytes".to_vec();
let encoded = anda_core::ByteBufB64(bytes.clone()).to_base64();
let file_encoded = anda_core::ByteBufB64(b"data uri attachment".to_vec()).to_base64();
let msg = anda_core::Message {
role: "user".into(),
content: vec![
"look".to_string().into(),
anda_core::ContentPart::InlineData {
mime_type: "image/png".into(),
data: anda_core::ByteBufB64(bytes.clone()),
},
anda_core::ContentPart::InlineData {
mime_type: "application/pdf".into(),
data: anda_core::ByteBufB64(bytes),
},
anda_core::ContentPart::FileData {
file_uri: format!("data:application/pdf;base64,{file_encoded}"),
mime_type: Some("application/pdf".into()),
},
],
..Default::default()
};
let mut raw = vec![serde_json::to_value(types::Content::from(msg)).unwrap()];
crate::model::raw::prune_inline_media(&mut raw);
let sent = serde_json::to_string(&raw).unwrap();
assert!(!sent.contains(&encoded), "{sent}");
assert!(!sent.contains(&file_encoded), "{sent}");
assert!(sent.contains("[inline image/png data omitted]"), "{sent}");
assert!(
sent.contains("[inline application/pdf data omitted]"),
"{sent}"
);
assert!(sent.contains("look"), "{sent}");
}
#[test]
fn prune_inline_media_keeps_gemini_function_responses_valid() {
let screenshot = anda_core::ByteBufB64(b"screenshot bytes".to_vec());
let encoded = screenshot.to_base64();
let msg = anda_core::Message {
role: "tool".into(),
content: vec![anda_core::ContentPart::ToolOutput {
name: "screenshot".into(),
output: anda_core::ToolPresentation {
text: "captured".into(),
media: vec![anda_core::ToolMedia {
mime_type: "image/png".into(),
data: screenshot,
}],
}
.into_output(),
is_error: None,
call_id: Some("call_1".into()),
remote_id: None,
}],
..Default::default()
};
let mut raw = vec![serde_json::to_value(types::Content::from(msg)).unwrap()];
assert!(serde_json::to_string(&raw).unwrap().contains(&encoded));
crate::model::raw::prune_inline_media(&mut raw);
let response = &raw[0]["parts"][0]["functionResponse"];
assert!(response.get("parts").is_none(), "{response}");
assert_eq!(
response["response"]["output"],
"captured\n[inline image/png data omitted]"
);
}
}
#[cfg(test)]
mod tool_presentation_compat_tests {
use super::*;
use crate::model::test_support::{no_proxy_client, recorded, spawn_mock_server};
use anda_core::{ByteBufB64, ContentPart, ToolMedia, ToolPresentation};
use http::{HeaderMap, StatusCode};
use serde_json::json;
#[tokio::test]
async fn text_fallback_preserves_neutral_media_for_later_model_replay() {
let presentation = ToolPresentation {
text: "result".into(),
media: vec![
ToolMedia {
mime_type: "image/png".into(),
data: ByteBufB64::from(vec![1, 2, 3]),
},
ToolMedia {
mime_type: "audio/wav".into(),
data: ByteBufB64::from(vec![4, 5, 6]),
},
],
};
let content = ContentPart::ToolOutput {
name: "echo".into(),
call_id: Some("call-1".into()),
is_error: None,
remote_id: None,
output: presentation.clone().into_output(),
};
let history = vec![
Message {
role: "user".into(),
content: vec!["Read the tool media.".to_string().into()],
..Default::default()
},
Message {
role: "assistant".into(),
content: vec![ContentPart::ToolCall {
name: "echo".into(),
args: json!({}),
call_id: Some("call-1".into()),
}],
..Default::default()
},
];
let raw_history = vec![
json!({"role":"user","parts":[{"text":"Read the tool media."}]}),
json!({"role":"model","parts":[{"functionCall":{"name":"echo","args":{},"id":"call-1"},"thoughtSignature":"provider-signature"}]}),
];
let body = serde_json::to_vec(&json!({
"candidates": [{
"content": {"role":"model", "parts":[{"text":"done"}]},
"finishReason":"STOP"
}]
}))
.unwrap();
let (endpoint, state) = spawn_mock_server(StatusCode::OK, HeaderMap::new(), body).await;
let client = Client::new_with_client("test-key", Some(endpoint), no_proxy_client());
for model in ["gemini-2.5-pro", "gemini-flash-latest"] {
let output = client
.completion_model(model)
.with_stream(false)
.completion(CompletionRequest {
raw_history: raw_history.clone(),
role: Some("tool".into()),
content: vec![content.clone()],
..Default::default()
})
.await
.unwrap();
let sent: Json = serde_json::from_slice(&recorded(&state).body).unwrap();
assert_eq!(&sent["contents"].as_array().unwrap()[..2], &raw_history);
let result = &sent["contents"][2]["parts"][0]["functionResponse"];
assert_eq!(result["id"], "call-1");
assert_eq!(result["response"]["output"], presentation.text_fallback());
assert!(result.get("parts").is_none());
assert_eq!(output.chat_history[0].role, "tool");
assert_eq!(output.chat_history[0].content, vec![content.clone()]);
let mut replay = history.clone();
replay.extend(output.chat_history);
client
.completion_model("gemini-3-flash")
.with_stream(false)
.completion(CompletionRequest {
chat_history: replay,
prompt: "Describe the image.".into(),
..Default::default()
})
.await
.unwrap();
let sent: Json = serde_json::from_slice(&recorded(&state).body).unwrap();
let result = &sent["contents"][2]["parts"][0]["functionResponse"];
assert_eq!(result["id"], "call-1");
assert_eq!(result["parts"][0]["inlineData"]["data"], "AQID");
assert_eq!(result["parts"][0]["inlineData"]["mimeType"], "image/png");
assert!(
result["response"]["output"]
.as_str()
.unwrap()
.contains("audio/wav")
);
assert!(
!result["response"]["output"]
.as_str()
.unwrap()
.contains("image/png")
);
}
}
}