use async_trait::async_trait;
use futures::stream::{Stream, StreamExt};
use reqwest::StatusCode;
use std::pin::Pin;
use std::sync::Arc;
use tokio::sync::RwLock;
use crate::error::{Error, Result};
use crate::providers::{
GenerationConfig, GenerationResponse, InferenceClient, StreamChunk, TraceCallback,
};
use crate::retry::{retry_with_backoff, RetryConfig};
use crate::types::{CacheMarker, ChatRole, Provider, TokenUsage};
use crate::utils::{
convert_messages_to_provider_format, parse_json_value_strict_str, ConversationMessage,
};
#[cfg(feature = "google-adc")]
#[derive(Clone)]
pub struct GoogleAnthropicClient {
model: String,
project_id: String,
region: String,
token_provider: Arc<RwLock<Option<Arc<dyn gcp_auth::TokenProvider>>>>,
thinking_budget: Option<u32>,
trace_callback: Option<TraceCallback>,
}
#[cfg(feature = "google-adc")]
impl GoogleAnthropicClient {
pub fn from_adc(model: impl Into<String>, region: impl Into<String>) -> Self {
let creds_path = std::env::var("GOOGLE_APPLICATION_CREDENTIALS")
.expect("GOOGLE_APPLICATION_CREDENTIALS environment variable must be set");
let creds_content = std::fs::read_to_string(&creds_path)
.unwrap_or_else(|_| panic!("Failed to read credentials file: {}", creds_path));
let creds_json: serde_json::Value = parse_json_value_strict_str(&creds_content)
.expect("Failed to parse credentials file as JSON");
let project_id = creds_json["project_id"]
.as_str()
.expect("Credentials file must contain project_id")
.to_string();
Self {
model: model.into(),
project_id,
region: region.into(),
token_provider: Arc::new(RwLock::new(None)),
thinking_budget: None,
trace_callback: None,
}
}
pub fn from_adc_with_project(
model: impl Into<String>,
project_id: impl Into<String>,
region: impl Into<String>,
) -> Self {
Self {
model: model.into(),
project_id: project_id.into(),
region: region.into(),
token_provider: Arc::new(RwLock::new(None)),
thinking_budget: None,
trace_callback: None,
}
}
pub fn with_thinking_budget(mut self, budget: u32) -> Self {
self.thinking_budget = Some(budget);
self
}
fn get_endpoint_url(&self) -> String {
format!(
"https://{}-aiplatform.googleapis.com/v1/projects/{}/locations/{}/publishers/anthropic/models/{}:streamRawPredict",
self.region, self.project_id, self.region, self.model
)
}
async fn get_token(&self) -> Result<String> {
let mut provider = self.token_provider.write().await;
if provider.is_none() {
let tp = gcp_auth::provider()
.await
.map_err(|e| Error::NonRetryable(format!("Failed to initialize ADC: {}", e)))?;
*provider = Some(tp);
}
let tp = provider.as_ref().unwrap();
let scopes = &["https://www.googleapis.com/auth/cloud-platform"];
let token = tp
.token(scopes)
.await
.map_err(|e| Error::NonRetryable(format!("Failed to get ADC token: {}", e)))?;
Ok(token.as_str().to_string())
}
fn build_request_body(
&self,
messages: &[ConversationMessage],
config: &GenerationConfig,
stream: bool,
) -> Result<serde_json::Value> {
let (system, messages) = self.extract_system_message(messages)?;
let formatted_messages =
convert_messages_to_provider_format(messages, Provider::Anthropic)?;
let formatted_messages = self.wrap_string_content(formatted_messages);
let _model = if config.model.is_empty() {
&self.model
} else {
&config.model
};
let mut request = serde_json::json!({
"anthropic_version": "vertex-2023-10-16",
"max_tokens": config.max_tokens.unwrap_or(4096),
"messages": formatted_messages,
"stream": stream,
});
if self.thinking_budget.is_none() {
if let Some(temp) = config.temperature {
request["temperature"] = serde_json::json!(temp);
}
if let Some(top_p) = config.top_p {
request["top_p"] = serde_json::json!(top_p);
}
}
if let Some(system) = system {
request["system"] = system;
}
if let Some(ref tools) = config.tools {
if !tools.is_empty() {
request["tools"] = serde_json::json!(tools);
request["tool_choice"] = serde_json::json!({
"type": "auto",
"disable_parallel_tool_use": false
});
}
}
if let Some(budget) = self.thinking_budget {
request["thinking"] = serde_json::json!({
"type": "enabled",
"budget_tokens": budget
});
let current_max = request["max_tokens"].as_u64().unwrap_or(4096);
request["max_tokens"] = serde_json::json!(current_max + budget as u64);
request["temperature"] = serde_json::json!(1.0);
if let Some(obj) = request.as_object_mut() {
obj.remove("top_p");
}
}
Ok(request)
}
fn extract_system_message<'a>(
&self,
messages: &'a [ConversationMessage],
) -> Result<(Option<serde_json::Value>, &'a [ConversationMessage])> {
if let Some(ConversationMessage::Chat(first)) = messages.first() {
if first.role == ChatRole::System {
let mut system_dict = serde_json::json!({
"type": "text",
"text": first.content
});
if first.cache_marker == Some(CacheMarker::Ephemeral) {
system_dict["cache_control"] = serde_json::json!({"type": "ephemeral"});
}
return Ok((Some(serde_json::json!([system_dict])), &messages[1..]));
}
}
Ok((None, messages))
}
fn wrap_string_content(&self, mut messages: Vec<serde_json::Value>) -> Vec<serde_json::Value> {
for msg in messages.iter_mut() {
if let Some(content) = msg.get("content") {
if let Some(text) = content.as_str() {
if !text.is_empty() {
msg["content"] = serde_json::json!([{
"type": "text",
"text": text
}]);
} else {
msg["content"] = serde_json::json!([]);
}
}
}
}
messages
}
fn extract_text_content(&self, content: &serde_json::Value) -> String {
if let Some(blocks) = content.as_array() {
for block in blocks {
if block["type"] == "text" {
if let Some(text) = block["text"].as_str() {
return text.to_string();
}
}
}
}
String::new()
}
fn extract_tool_calls(&self, content: &serde_json::Value) -> Vec<serde_json::Value> {
let mut tool_calls = Vec::new();
if let Some(blocks) = content.as_array() {
for block in blocks {
if block["type"] == "tool_use" {
tool_calls.push(block.clone());
}
}
}
tool_calls
}
fn parse_usage(&self, usage: &serde_json::Value) -> TokenUsage {
TokenUsage {
input_tokens: usage["input_tokens"].as_u64().unwrap_or(0),
output_tokens: usage["output_tokens"].as_u64().unwrap_or(0),
cached_tokens: usage["cache_read_input_tokens"].as_u64().unwrap_or(0),
}
}
async fn make_request(
&self,
body: serde_json::Value,
timeout: Option<std::time::Duration>,
) -> Result<reqwest::Response> {
let token = self.get_token().await?;
let client = reqwest::Client::new();
let response = client
.post(self.get_endpoint_url())
.timeout(timeout.unwrap_or(std::time::Duration::from_secs(300)))
.header("Authorization", format!("Bearer {}", token))
.header("Content-Type", "application/json")
.json(&body)
.send()
.await?;
Ok(response)
}
fn handle_error_response(&self, status: StatusCode, body: String) -> Error {
match status {
StatusCode::BAD_REQUEST | StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => {
Error::NonRetryable(format!("{}: {}", status, body))
}
StatusCode::TOO_MANY_REQUESTS => Error::Inference(format!("Rate limited: {}", body)),
StatusCode::INTERNAL_SERVER_ERROR
| StatusCode::BAD_GATEWAY
| StatusCode::SERVICE_UNAVAILABLE
| StatusCode::GATEWAY_TIMEOUT => Error::Inference(format!("{}: {}", status, body)),
_ => Error::Inference(format!("{}: {}", status, body)),
}
}
#[cfg(feature = "google-adc")]
async fn make_request_with_retry(
&self,
body: serde_json::Value,
timeout: Option<std::time::Duration>,
) -> Result<serde_json::Value> {
let config = RetryConfig::default();
retry_with_backoff(config, || async {
let response = self.make_request(body.clone(), timeout).await?;
let status = response.status();
if !status.is_success() {
let error_body = response.text().await.unwrap_or_default();
return Err(self.handle_error_response(status, error_body));
}
response.json().await.map_err(Error::from)
})
.await
}
}
#[cfg(feature = "google-adc")]
#[async_trait]
impl InferenceClient for GoogleAnthropicClient {
async fn get_generation(
&self,
messages: &[ConversationMessage],
config: &GenerationConfig,
) -> Result<GenerationResponse> {
let request_body = self.build_request_body(messages, config, false)?;
let body = self
.make_request_with_retry(request_body, config.timeout)
.await?;
let usage = self.parse_usage(&body["usage"]);
let content_value = &body["content"];
let text_content = self.extract_text_content(content_value);
let tool_calls = self.extract_tool_calls(content_value);
let has_tools = config.tools.is_some() && !config.tools.as_ref().unwrap().is_empty();
let has_tool_calls = !tool_calls.is_empty();
Ok(GenerationResponse {
content: text_content,
reasoning: None,
tool_calls,
reasoning_segments: Vec::new(),
usage,
provider_cost_dollars: None,
raw: if has_tools || has_tool_calls {
Some(body)
} else {
None
},
})
}
async fn connect_and_listen(
&self,
messages: &[ConversationMessage],
config: &GenerationConfig,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamChunk>> + Send>>> {
use crate::providers::anthropic_streaming::{parse_anthropic_chunk, ToolCallAccumulator};
use crate::utils::parse_sse_stream;
let request_body = self.build_request_body(messages, config, true)?;
let timeout = config.timeout;
let retry_config = RetryConfig::default();
let response = retry_with_backoff(retry_config, || async {
let resp = self.make_request(request_body.clone(), timeout).await?;
let status = resp.status();
if !status.is_success() {
let error_body = resp.text().await.unwrap_or_default();
return Err(self.handle_error_response(status, error_body));
}
Ok(resp)
})
.await?;
let has_tools = config.tools.is_some() && !config.tools.as_ref().unwrap().is_empty();
let sse_stream = parse_sse_stream(response);
let chunk_stream = sse_stream.scan(
ToolCallAccumulator::new(),
move |accumulator, sse_result| {
let sse_json = match sse_result {
Ok(json) => json,
Err(e) => return futures::future::ready(Some(vec![Err(e)])),
};
let chunks = parse_anthropic_chunk(&sse_json, accumulator, has_tools);
futures::future::ready(Some(chunks.into_iter().map(Ok).collect()))
},
);
Ok(Box::pin(chunk_stream.flat_map(futures::stream::iter)))
}
fn provider(&self) -> Provider {
Provider::GoogleAnthropic
}
fn set_trace_callback(&mut self, callback: TraceCallback) {
self.trace_callback = Some(callback);
}
}
#[cfg(all(test, feature = "google-adc"))]
mod tests {
use super::*;
#[test]
fn test_endpoint_url() {
let client = GoogleAnthropicClient::from_adc_with_project(
"claude-3-5-sonnet-v2@20241022",
"my-project",
"us-east5",
);
assert_eq!(
client.get_endpoint_url(),
"https://us-east5-aiplatform.googleapis.com/v1/projects/my-project/locations/us-east5/publishers/anthropic/models/claude-3-5-sonnet-v2@20241022:streamRawPredict"
);
}
#[test]
fn test_with_thinking_budget() {
let client = GoogleAnthropicClient::from_adc_with_project(
"claude-3-5-sonnet-v2@20241022",
"my-project",
"us-east5",
)
.with_thinking_budget(1024);
assert_eq!(client.thinking_budget, Some(1024));
}
}