use async_trait::async_trait;
use mermaid_domain::ChatRequest;
use mermaid_model::models::adapters::gemini::GeminiAdapter;
use mermaid_model::models::{Model, ModelConfig, ModelError, Result};
use super::super::ctx::{FinalResponse, StreamContext, StreamEvent};
use super::{ContextSizing, ModelProvider, RejectionCache, resolve_limits_cached};
use mermaid_model::models::ModelCapabilities;
pub const DEFAULT_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta";
pub const DEFAULT_API_KEY_ENV: &str = "GOOGLE_API_KEY";
pub const LEGACY_API_KEY_ENV: &str = "GEMINI_API_KEY";
pub struct GeminiProvider {
adapter: GeminiAdapter,
capabilities: ModelCapabilities,
rejections: RejectionCache,
}
impl GeminiProvider {
pub fn new(api_key: String, model_name: String, base_url: String) -> Result<Self> {
let adapter = GeminiAdapter::new(api_key, model_name, base_url)?;
let capabilities = adapter.capabilities().clone();
Ok(Self {
adapter,
capabilities,
rejections: RejectionCache::default(),
})
}
}
#[async_trait]
impl ModelProvider for GeminiProvider {
fn capabilities(&self) -> &ModelCapabilities {
&self.capabilities
}
async fn resolve_context_window(&self, request: &ChatRequest) -> ContextSizing {
let _ = request;
let model = Model::name(&self.adapter).to_string();
let limits =
resolve_limits_cached("gemini", &model, || self.adapter.fetch_model_limits()).await;
let window = limits.as_ref().and_then(|l| l.max_context_tokens);
ContextSizing {
model_max: window,
effective: window,
source: None,
max_output: limits.as_ref().and_then(|l| l.max_output_tokens),
compacts_natively: false,
}
}
async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
let config = ModelConfig::from(&request);
let chat_fut = self
.adapter
.chat(&request.messages, &config, Some(ctx.sink.clone()));
let learned_model = Model::name(&self.adapter).to_string();
let memory = self.adapter.param_memory();
self.rejections.seed("gemini", &learned_model, memory).await;
let response = tokio::select! {
biased;
_ = ctx.token.cancelled() => {
return Err(ModelError::Cancelled);
},
r = chat_fut => {
self.rejections.persist("gemini", &learned_model, memory).await;
r?
},
};
let usage = response.usage.clone();
let stop_reason = response.stop_reason.clone();
let _ = ctx
.sink
.send(StreamEvent::Done {
usage: usage.clone(),
provider_continuation: None,
stop_reason: stop_reason.clone(),
})
.await;
Ok(FinalResponse {
usage,
provider_continuation: None,
tool_calls: response.tool_calls.unwrap_or_default(),
stop_reason,
})
}
}