use async_trait::async_trait;
use serde_json::{Value, json};
use crate::core::Secret;
use super::wire::{RESPOND_TOOL, classify_status, classify_transport, structured};
use super::{
Completion, ModelError, ModelId, ModelProvider, ReasoningEffort, Request, SchemaMode, Usage,
gemini_stream, sse,
};
pub(crate) const PROVIDER: &str = "gemini";
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
#[non_exhaustive]
pub enum HarmCategory {
Harassment,
HateSpeech,
SexuallyExplicit,
DangerousContent,
CivicIntegrity,
Jailbreak,
}
impl HarmCategory {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Harassment => "HARM_CATEGORY_HARASSMENT",
Self::HateSpeech => "HARM_CATEGORY_HATE_SPEECH",
Self::SexuallyExplicit => "HARM_CATEGORY_SEXUALLY_EXPLICIT",
Self::DangerousContent => "HARM_CATEGORY_DANGEROUS_CONTENT",
Self::CivicIntegrity => "HARM_CATEGORY_CIVIC_INTEGRITY",
Self::Jailbreak => "HARM_CATEGORY_JAILBREAK",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
#[non_exhaustive]
pub enum HarmBlockThreshold {
LowAndAbove,
MediumAndAbove,
OnlyHigh,
None,
}
impl HarmBlockThreshold {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::LowAndAbove => "BLOCK_LOW_AND_ABOVE",
Self::MediumAndAbove => "BLOCK_MEDIUM_AND_ABOVE",
Self::OnlyHigh => "BLOCK_ONLY_HIGH",
Self::None => "BLOCK_NONE",
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct SafetySettings {
thresholds: std::collections::BTreeMap<HarmCategory, HarmBlockThreshold>,
}
impl SafetySettings {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn block(mut self, category: HarmCategory, threshold: HarmBlockThreshold) -> Self {
self.thresholds.insert(category, threshold);
self
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.thresholds.is_empty()
}
fn wire(&self) -> Value {
Value::Array(
self.thresholds
.iter()
.map(|(category, threshold)| {
json!({ "category": category.as_str(), "threshold": threshold.as_str() })
})
.collect(),
)
}
fn profile(&self) -> Value {
self.wire()
}
}
pub struct Gemini {
http: reqwest::Client,
key: Secret,
base: String,
version: String,
default_schema_mode: SchemaMode,
schema_modes: std::collections::BTreeMap<String, SchemaMode>,
stream: bool,
egress: Option<crate::core::Egress>,
timeout: std::time::Duration,
safety: SafetySettings,
}
impl std::fmt::Debug for Gemini {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Gemini")
.field("base", &self.base)
.field("version", &self.version)
.field("key", &"<redacted>")
.finish_non_exhaustive()
}
}
impl Gemini {
pub const DEFAULT_TIMEOUT: std::time::Duration = std::time::Duration::from_mins(5);
pub const DEFAULT_BASE: &'static str = "https://generativelanguage.googleapis.com";
pub const VERSION: &'static str = "v1beta";
pub fn new(key: impl Into<String>) -> Result<Self, ModelError> {
let http = crate::netguard::guarded_client(crate::netguard::Reach::Configured)
.build()
.map_err(|e| ModelError::Unreachable {
model: ModelId::new(PROVIDER, "*"),
detail: format!("could not build an HTTP client: {e}"),
})?;
Ok(Self {
http,
key: Secret::new(key),
base: Self::DEFAULT_BASE.to_owned(),
version: Self::VERSION.to_owned(),
default_schema_mode: SchemaMode::Native,
schema_modes: std::collections::BTreeMap::new(),
stream: true,
egress: None,
timeout: Self::DEFAULT_TIMEOUT,
safety: SafetySettings::new(),
})
}
pub fn from_env() -> Result<Self, ModelError> {
let key = std::env::var("GEMINI_API_KEY")
.or_else(|_| std::env::var("GOOGLE_API_KEY"))
.map_err(|_| ModelError::Refused {
model: ModelId::new(PROVIDER, "*"),
detail: "neither GEMINI_API_KEY nor GOOGLE_API_KEY is set".to_owned(),
})?;
Self::new(key)
}
#[must_use]
pub fn base(mut self, base: impl Into<String>) -> Self {
let mut base = base.into();
while base.ends_with('/') {
base.pop();
}
self.base = base;
self
}
#[must_use]
pub const fn timeout(mut self, timeout: std::time::Duration) -> Self {
self.timeout = timeout;
self
}
#[must_use]
pub fn structured_via(mut self, mode: SchemaMode) -> Self {
self.default_schema_mode = mode;
self
}
#[must_use]
pub fn structured_via_for(mut self, model: impl Into<String>, mode: SchemaMode) -> Self {
self.schema_modes.insert(model.into(), mode);
self
}
#[must_use]
pub fn safety(mut self, safety: SafetySettings) -> Self {
self.safety = safety;
self
}
#[must_use]
pub fn egress(mut self, egress: crate::core::Egress) -> Self {
self.egress = Some(egress);
self
}
#[must_use]
pub const fn buffered(mut self) -> Self {
self.stream = false;
self
}
fn mode_for(&self, model: &ModelId) -> SchemaMode {
self.schema_modes
.get(&model.model)
.copied()
.unwrap_or(self.default_schema_mode)
}
fn check_egress(&self, model: &ModelId) -> Result<(), ModelError> {
let Some(egress) = &self.egress else {
return Ok(());
};
let host = reqwest::Url::parse(&self.base)
.ok()
.and_then(|u| u.host_str().map(ToOwned::to_owned));
egress
.permits(host.as_deref())
.map_err(|e| ModelError::Egress {
model: model.clone(),
detail: e.to_string(),
})
}
fn refused(model: &ModelId, detail: impl Into<String>) -> ModelError {
ModelError::Refused {
model: model.clone(),
detail: detail.into(),
}
}
fn thinking_config(model: &ModelId, effort: ReasoningEffort) -> Result<Value, ModelError> {
let level = match effort {
ReasoningEffort::Minimal => "minimal",
ReasoningEffort::Low => "low",
ReasoningEffort::Medium => "medium",
ReasoningEffort::High => "high",
ReasoningEffort::None | ReasoningEffort::XHigh | ReasoningEffort::Max => {
return Err(Self::refused(
model,
format!(
"Gemini has no thinking level for reasoning effort '{}' — it names \
minimal, low, medium and high, and thinking cannot be switched off \
on the Gemini 3 models",
effort.as_str()
),
));
}
};
Ok(json!({ "thinkingLevel": level }))
}
fn contents(prompt: &Value) -> Value {
match prompt {
Value::String(text) => json!([{ "role": "user", "parts": [{ "text": text }] }]),
Value::Array(_) => prompt.clone(),
other => crate::model::prompt_envelope(other, "messages")
.cloned()
.unwrap_or_else(|| {
let mut rest = other.clone();
if let Some(map) = rest.as_object_mut() {
map.remove("system");
}
json!([{ "role": "user", "parts": [{ "text": rest.to_string() }] }])
}),
}
}
fn system_instruction(prompt: &Value) -> Option<Value> {
let system = prompt.get("system").filter(|s| !s.is_null())?;
Some(match system {
Value::String(text) => json!({ "parts": [{ "text": text }] }),
other => other.clone(),
})
}
fn body(&self, model: &ModelId, request: &Request<'_>) -> Result<Value, ModelError> {
let Request {
prompt,
max_output_tokens,
reasoning_effort,
schema,
tools,
exchanges,
continuation,
..
} = request;
let mut contents = Self::contents(prompt);
Self::append_tool_turns(&mut contents, exchanges, *continuation, model)?;
let mut generation_config = json!({ "maxOutputTokens": max_output_tokens });
if let Some(effort) = reasoning_effort {
generation_config["thinkingConfig"] = Self::thinking_config(model, *effort)?;
}
let mut body = json!({ "contents": contents });
if !self.safety.is_empty() {
body["safetySettings"] = self.safety.wire();
}
if let Some(system) = Self::system_instruction(prompt) {
body["systemInstruction"] = system;
}
let mode = self.mode_for(model);
let mut declarations: Vec<Value> = tools
.iter()
.map(|t| {
json!({
"name": t.name,
"description": t.description,
"parametersJsonSchema": t.parameters,
})
})
.collect();
if let Some(schema) = schema {
match mode {
SchemaMode::Native => {
generation_config["responseMimeType"] = json!("application/json");
generation_config["responseJsonSchema"] = (*schema).clone();
}
SchemaMode::ForcedTool => {
if !tools.is_empty() {
return Err(Self::refused(
model,
"forced-tool structured output cannot be combined with declared \
tools: the model would be offered a choice between answering and \
calling one. Use SchemaMode::Native, which Gemini enforces during \
generation",
));
}
declarations.push(json!({
"name": RESPOND_TOOL,
"description": "Return the answer in the required shape.",
"parametersJsonSchema": (*schema).clone(),
}));
body["toolConfig"] = json!({
"functionCallingConfig": {
"mode": "ANY",
"allowedFunctionNames": [RESPOND_TOOL],
}
});
}
}
}
if !declarations.is_empty() {
body["tools"] = json!([{ "functionDeclarations": declarations }]);
}
body["generationConfig"] = generation_config;
Ok(body)
}
fn append_tool_turns(
contents: &mut Value,
exchanges: &[super::ToolExchange],
continuation: Option<&super::ProviderContinuation>,
model: &ModelId,
) -> Result<(), ModelError> {
if exchanges.is_empty() {
if continuation.is_some() {
return Err(Self::refused(
model,
"a continuation without tool exchanges has no request to follow",
));
}
return Ok(());
}
let Some(array) = contents.as_array_mut() else {
return Err(Self::refused(
model,
"the prompt did not assemble into a `contents` array",
));
};
match continuation {
Some(state) if state.provider == PROVIDER => match state.state.as_array() {
Some(turns) => array.extend(turns.iter().cloned()),
None => {
return Err(Self::refused(
model,
"the continuation was not a Gemini contents array",
));
}
},
Some(other) => {
return Err(Self::refused(
model,
format!(
"the continuation was issued by '{}' and this is the Gemini driver — \
provider state is opaque and is never valid across providers",
other.provider
),
));
}
None => array.push(json!({
"role": "model",
"parts": exchanges
.iter()
.map(|e| json!({
"functionCall": { "name": e.call.name, "args": e.call.arguments }
}))
.collect::<Vec<_>>(),
})),
}
array.push(Self::tool_responses(exchanges));
Ok(())
}
fn tool_responses(exchanges: &[super::ToolExchange]) -> Value {
json!({
"role": "user",
"parts": exchanges
.iter()
.map(|e| {
let mut response = json!({ "name": e.call.name, "response": {
"output": e.output,
}});
if e.failed {
response["response"] = json!({ "error": e.output });
}
json!({ "functionResponse": response })
})
.collect::<Vec<_>>(),
})
}
fn accumulate_continuation(
completion: &mut Completion,
prior: Option<&super::ProviderContinuation>,
exchanges: &[super::ToolExchange],
) {
let Some(current) = completion.continuation.as_mut() else {
return;
};
let mut state = prior
.and_then(|value| value.state.as_array())
.cloned()
.unwrap_or_default();
if !exchanges.is_empty() {
state.push(Self::tool_responses(exchanges));
}
if let Some(turns) = current.state.as_array() {
state.extend(turns.iter().cloned());
}
current.state = Value::Array(state);
}
fn interpret(
&self,
parsed: &Value,
model: &ModelId,
schema: Option<&Value>,
) -> Result<Completion, ModelError> {
let usage = Self::usage(parsed);
let Some(candidate) = parsed.get("candidates").and_then(|c| c.get(0)) else {
let detail = parsed
.get("promptFeedback")
.and_then(|f| f.get("blockReason"))
.and_then(Value::as_str)
.map_or_else(
|| "the response carried no candidates".to_owned(),
|reason| format!("the prompt was blocked before generating: {reason}"),
);
return Err(Self::refused(model, detail));
};
let finish = candidate
.get("finishReason")
.and_then(Value::as_str)
.map(ToOwned::to_owned);
let truncated = finish.as_deref() == Some("MAX_TOKENS");
if let Some(reason) = finish.as_deref()
&& matches!(
reason,
"SAFETY" | "RECITATION" | "PROHIBITED_CONTENT" | "BLOCKLIST" | "SPII"
)
{
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: format!("generation stopped: {reason}"),
});
}
let content = candidate.get("content").cloned().unwrap_or(Value::Null);
let parts = content
.get("parts")
.and_then(Value::as_array)
.cloned()
.unwrap_or_default();
let Scanned {
text,
calls,
forced,
} = Self::scan(&parts);
let emulating = schema.is_some() && self.mode_for(model) == SchemaMode::ForcedTool;
if text.is_empty() && calls.is_empty() && forced.is_none() && !truncated {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: format!("the answer carried no content (finishReason {finish:?})"),
});
}
let (text, structured_value) = if emulating {
let Some(arguments) = forced else {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: "a tool call was forced and the answer carried none — the model \
did not honour the function-calling config"
.to_owned(),
});
};
let raw = arguments.to_string();
let parsed_schema = structured(schema, &raw, &[], model, usage)?;
(raw, parsed_schema)
} else {
let parsed_schema = structured(schema, &text, &calls, model, usage)?;
(text, parsed_schema)
};
let continuation = (!calls.is_empty() && content.is_object())
.then(|| super::ProviderContinuation::new(PROVIDER, json!([content.clone()])));
Ok(Completion {
structured: structured_value,
tool_calls: calls,
text,
usage,
stop_reason: finish,
truncated,
continuation,
})
}
fn usage(parsed: &Value) -> Usage {
let count = |key: &str| {
parsed
.get("usageMetadata")
.and_then(|u| u.get(key))
.and_then(Value::as_u64)
.unwrap_or_default()
};
Usage {
input_tokens: count("promptTokenCount"),
output_tokens: count("candidatesTokenCount") + count("thoughtsTokenCount"),
cache_read_tokens: count("cachedContentTokenCount"),
cache_write_tokens: 0,
minor_units: 0,
}
}
fn url(&self, model: &ModelId) -> String {
let method = if self.stream {
"streamGenerateContent?alt=sse"
} else {
"generateContent"
};
format!(
"{}/{}/models/{}:{method}",
self.base, self.version, model.model
)
}
async fn read_buffered(
&self,
response: reqwest::Response,
model: &ModelId,
schema: Option<&Value>,
) -> Result<Completion, ModelError> {
let body = crate::netguard::intake::read(response, crate::netguard::intake::ANSWER)
.await
.map_err(|e| super::wire::classify_intake(model, Usage::default(), &e))?;
let parsed: Value = serde_json::from_slice(&body).map_err(|e| ModelError::Unusable {
model: model.clone(),
usage: Usage::default(),
detail: format!("the response did not parse: {e}"),
})?;
self.interpret(&parsed, model, schema)
}
async fn read_streamed(
&self,
response: reqwest::Response,
model: &ModelId,
schema: Option<&Value>,
observer: Option<(&dyn super::ModelStreamObserver, &crate::core::Label)>,
) -> Result<Completion, ModelError> {
use futures_util::StreamExt;
let mut decoder = sse::Decoder::new();
let mut acc = gemini_stream::Accumulator::new();
let mut body = response.bytes_stream();
let mut meter = crate::netguard::intake::Meter::new(crate::netguard::intake::ANSWER);
while let Some(chunk) = body.next().await {
let chunk = match chunk {
Ok(chunk) => chunk,
Err(e) => return Err(severed(model, &acc, &e.to_string())),
};
if let Err(e) = meter.charge(chunk.len()) {
return Err(super::wire::classify_intake(model, Usage::default(), &e));
}
let events = decoder
.push(&chunk)
.map_err(|error| severed(model, &acc, &error.to_string()))?;
for event in events {
if let Some(delta) = acc.push(&event.data)
&& let Some((observer, label)) = observer
{
observer.event(crate::core::Tainted::with_label(
super::ModelStreamEvent::TextDelta(delta),
label.clone(),
));
}
}
}
if !acc.done() {
if acc.prompt_blocked() {
return self.interpret(&acc.into_response(), model, schema);
}
return Err(severed(
model,
&acc,
"the stream ended before the model said why it stopped",
));
}
let completion = self.interpret(&acc.into_response(), model, schema)?;
if let Some((observer, label)) = observer {
observer.event(crate::core::Tainted::with_label(
super::ModelStreamEvent::Usage(completion.usage),
label.clone(),
));
}
Ok(completion)
}
}
struct Scanned {
text: String,
calls: Vec<super::ToolCall>,
forced: Option<Value>,
}
impl Gemini {
fn scan(parts: &[Value]) -> Scanned {
let mut text = String::new();
let mut calls = Vec::new();
let mut forced: Option<Value> = None;
for (index, part) in parts.iter().enumerate() {
if part.get("thought").and_then(Value::as_bool) == Some(true) {
continue;
}
if let Some(call) = part.get("functionCall") {
let name = call
.get("name")
.and_then(Value::as_str)
.unwrap_or_default()
.to_owned();
let arguments = call.get("args").cloned().unwrap_or_else(|| json!({}));
if name == RESPOND_TOOL {
forced = Some(arguments);
continue;
}
calls.push(super::ToolCall {
id: call
.get("id")
.and_then(Value::as_str)
.map_or_else(|| format!("{name}-{index}"), ToOwned::to_owned),
name,
arguments,
});
continue;
}
if let Some(chunk) = part.get("text").and_then(Value::as_str) {
text.push_str(chunk);
}
}
Scanned {
text,
calls,
forced,
}
}
}
fn severed(model: &ModelId, acc: &gemini_stream::Accumulator, detail: &str) -> ModelError {
if let Some(envelope) = acc.usage_envelope() {
return ModelError::Interrupted {
model: model.clone(),
usage: Gemini::usage(&envelope),
detail: detail.to_owned(),
};
}
if acc.generated() {
return ModelError::Unaccounted {
model: model.clone(),
detail: detail.to_owned(),
};
}
ModelError::Unavailable {
model: model.clone(),
detail: detail.to_owned(),
}
}
fn with_retry_info(error: ModelError, body: &str) -> ModelError {
match error {
ModelError::RateLimited {
model,
detail,
retry_after: None,
} => ModelError::RateLimited {
model,
detail,
retry_after: retry_info_seconds(body),
},
other => other,
}
}
fn retry_info_seconds(body: &str) -> Option<u64> {
let parsed: Value = serde_json::from_str(body).ok()?;
let details = parsed.get("error")?.get("details")?.as_array()?;
let delay = details
.iter()
.find(|d| {
d.get("@type").and_then(Value::as_str)
== Some("type.googleapis.com/google.rpc.RetryInfo")
})?
.get("retryDelay")?
.as_str()?;
let seconds: u64 = delay.strip_suffix('s')?.split('.').next()?.parse().ok()?;
(seconds > 0).then_some(seconds)
}
#[async_trait]
impl ModelProvider for Gemini {
fn request_profile(&self, model: &ModelId) -> Value {
json!({
"driver": "google-gemini-generatecontent/v1",
"base": self.base,
"api_version": self.version,
"stream": self.stream,
"schema_mode": match self.mode_for(model) {
SchemaMode::Native => "native",
SchemaMode::ForcedTool => "forced-tool",
},
"timeout_ms": u64::try_from(self.timeout.as_millis()).unwrap_or(u64::MAX),
"safety": (!self.safety.is_empty()).then(|| self.safety.profile()),
})
}
async fn complete(&self, request: Request<'_>) -> Result<Completion, ModelError> {
let model = request.model;
super::refuse_provider_side_media(request.prompt, model)?;
super::refuse_in_thread_instructions(request.prompt, model)?;
self.check_egress(model)?;
let body = self.body(model, &request)?;
let response = self
.http
.post(self.url(model))
.header("x-goog-api-key", self.key.expose())
.timeout(self.timeout)
.json(&body)
.send()
.await
.map_err(|e| classify_transport(model, &e))?;
let status = response.status();
if !status.is_success() {
let headers = response.headers().clone();
let text =
crate::netguard::intake::read_text(response, crate::netguard::intake::METADATA)
.await
.unwrap_or_default();
return Err(with_retry_info(
classify_status(model, status.as_u16(), &headers, &text),
&text,
));
}
let mut completion = if self.stream {
self.read_streamed(response, model, request.schema, request.stream)
.await?
} else {
self.read_buffered(response, model, request.schema).await?
};
if !self.stream
&& let Some((observer, label)) = request.stream
{
observer.event(crate::core::Tainted::with_label(
super::ModelStreamEvent::Usage(completion.usage),
label.clone(),
));
}
Self::accumulate_continuation(&mut completion, request.continuation, request.exchanges);
Ok(completion)
}
}