feat: Azure credential chain logging (#2413)
Co-authored-by: Michael Paquette <mpaquette@pax8.com> Co-authored-by: Alice Hau <ahau@squareup.com>
This commit is contained in:
@@ -1,9 +1,12 @@
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use reqwest::Client;
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
use std::time::Duration;
|
||||
use tokio::time::sleep;
|
||||
|
||||
use super::azureauth::AzureAuth;
|
||||
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::errors::ProviderError;
|
||||
use super::formats::openai::{create_request, get_usage, response_to_message};
|
||||
@@ -18,17 +21,36 @@ pub const AZURE_DOC_URL: &str =
|
||||
pub const AZURE_DEFAULT_API_VERSION: &str = "2024-10-21";
|
||||
pub const AZURE_OPENAI_KNOWN_MODELS: &[&str] = &["gpt-4o", "gpt-4o-mini", "gpt-4"];
|
||||
|
||||
#[derive(Debug, serde::Serialize)]
|
||||
// Default retry configuration
|
||||
const DEFAULT_MAX_RETRIES: usize = 5;
|
||||
const DEFAULT_INITIAL_RETRY_INTERVAL_MS: u64 = 1000; // Start with 1 second
|
||||
const DEFAULT_MAX_RETRY_INTERVAL_MS: u64 = 32000; // Max 32 seconds
|
||||
const DEFAULT_BACKOFF_MULTIPLIER: f64 = 2.0;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct AzureProvider {
|
||||
#[serde(skip)]
|
||||
client: Client,
|
||||
auth: AzureAuth,
|
||||
endpoint: String,
|
||||
api_key: String,
|
||||
deployment_name: String,
|
||||
api_version: String,
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl Serialize for AzureProvider {
|
||||
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
S: serde::Serializer,
|
||||
{
|
||||
use serde::ser::SerializeStruct;
|
||||
let mut state = serializer.serialize_struct("AzureProvider", 3)?;
|
||||
state.serialize_field("endpoint", &self.endpoint)?;
|
||||
state.serialize_field("deployment_name", &self.deployment_name)?;
|
||||
state.serialize_field("api_version", &self.api_version)?;
|
||||
state.end()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for AzureProvider {
|
||||
fn default() -> Self {
|
||||
let model = ModelConfig::new(AzureProvider::metadata().default_model);
|
||||
@@ -39,13 +61,16 @@ impl Default for AzureProvider {
|
||||
impl AzureProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let api_key: String = config.get_secret("AZURE_OPENAI_API_KEY")?;
|
||||
let endpoint: String = config.get_param("AZURE_OPENAI_ENDPOINT")?;
|
||||
let deployment_name: String = config.get_param("AZURE_OPENAI_DEPLOYMENT_NAME")?;
|
||||
let api_version: String = config
|
||||
.get_param("AZURE_OPENAI_API_VERSION")
|
||||
.unwrap_or_else(|_| AZURE_DEFAULT_API_VERSION.to_string());
|
||||
|
||||
// Try to get API key first, if not found use Azure credential chain
|
||||
let api_key = config.get_secret("AZURE_OPENAI_API_KEY").ok();
|
||||
let auth = AzureAuth::new(api_key)?;
|
||||
|
||||
let client = Client::builder()
|
||||
.timeout(Duration::from_secs(600))
|
||||
.build()?;
|
||||
@@ -53,7 +78,7 @@ impl AzureProvider {
|
||||
Ok(Self {
|
||||
client,
|
||||
endpoint,
|
||||
api_key,
|
||||
auth,
|
||||
deployment_name,
|
||||
api_version,
|
||||
model,
|
||||
@@ -81,15 +106,110 @@ impl AzureProvider {
|
||||
base_url.set_path(&new_path);
|
||||
base_url.set_query(Some(&format!("api-version={}", self.api_version)));
|
||||
|
||||
let response: reqwest::Response = self
|
||||
.client
|
||||
.post(base_url)
|
||||
.header("api-key", &self.api_key)
|
||||
.json(&payload)
|
||||
.send()
|
||||
.await?;
|
||||
let mut attempts = 0;
|
||||
let mut last_error = None;
|
||||
let mut current_delay = DEFAULT_INITIAL_RETRY_INTERVAL_MS;
|
||||
|
||||
handle_response_openai_compat(response).await
|
||||
loop {
|
||||
// Check if we've exceeded max retries
|
||||
if attempts > DEFAULT_MAX_RETRIES {
|
||||
let error_msg = format!(
|
||||
"Exceeded maximum retry attempts ({}) for rate limiting",
|
||||
DEFAULT_MAX_RETRIES
|
||||
);
|
||||
tracing::error!("{}", error_msg);
|
||||
return Err(last_error.unwrap_or(ProviderError::RateLimitExceeded(error_msg)));
|
||||
}
|
||||
|
||||
// Get a fresh auth token for each attempt
|
||||
let auth_token = self.auth.get_token().await.map_err(|e| {
|
||||
tracing::error!("Authentication error: {:?}", e);
|
||||
ProviderError::RequestFailed(format!("Failed to get authentication token: {}", e))
|
||||
})?;
|
||||
|
||||
let mut request_builder = self.client.post(base_url.clone());
|
||||
let token_value = auth_token.token_value.clone();
|
||||
|
||||
// Set the correct header based on authentication type
|
||||
match self.auth.credential_type() {
|
||||
super::azureauth::AzureCredentials::ApiKey(_) => {
|
||||
request_builder = request_builder.header("api-key", token_value.clone());
|
||||
}
|
||||
super::azureauth::AzureCredentials::DefaultCredential => {
|
||||
request_builder = request_builder
|
||||
.header("Authorization", format!("Bearer {}", token_value.clone()));
|
||||
}
|
||||
}
|
||||
|
||||
let response_result = request_builder.json(&payload).send().await;
|
||||
|
||||
match response_result {
|
||||
Ok(response) => match handle_response_openai_compat(response).await {
|
||||
Ok(result) => {
|
||||
return Ok(result);
|
||||
}
|
||||
Err(ProviderError::RateLimitExceeded(msg)) => {
|
||||
attempts += 1;
|
||||
last_error = Some(ProviderError::RateLimitExceeded(msg.clone()));
|
||||
|
||||
let retry_after =
|
||||
if let Some(secs) = msg.to_lowercase().find("try again in ") {
|
||||
msg[secs..]
|
||||
.split_whitespace()
|
||||
.nth(3)
|
||||
.and_then(|s| s.parse::<u64>().ok())
|
||||
.unwrap_or(0)
|
||||
} else {
|
||||
0
|
||||
};
|
||||
|
||||
let delay = if retry_after > 0 {
|
||||
Duration::from_secs(retry_after)
|
||||
} else {
|
||||
let delay = current_delay.min(DEFAULT_MAX_RETRY_INTERVAL_MS);
|
||||
current_delay =
|
||||
(current_delay as f64 * DEFAULT_BACKOFF_MULTIPLIER) as u64;
|
||||
Duration::from_millis(delay)
|
||||
};
|
||||
|
||||
sleep(delay).await;
|
||||
continue;
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(
|
||||
"Error response from Azure OpenAI (attempt {}): {:?}",
|
||||
attempts + 1,
|
||||
e
|
||||
);
|
||||
return Err(e);
|
||||
}
|
||||
},
|
||||
Err(e) => {
|
||||
tracing::error!(
|
||||
"Request failed (attempt {}): {:?}\nIs timeout: {}\nIs connect: {}\nIs request: {}",
|
||||
attempts + 1,
|
||||
e,
|
||||
e.is_timeout(),
|
||||
e.is_connect(),
|
||||
e.is_request(),
|
||||
);
|
||||
|
||||
// For timeout errors, we should retry
|
||||
if e.is_timeout() {
|
||||
attempts += 1;
|
||||
let delay = current_delay.min(DEFAULT_MAX_RETRY_INTERVAL_MS);
|
||||
current_delay = (current_delay as f64 * DEFAULT_BACKOFF_MULTIPLIER) as u64;
|
||||
sleep(Duration::from_millis(delay)).await;
|
||||
continue;
|
||||
}
|
||||
|
||||
return Err(ProviderError::RequestFailed(format!(
|
||||
"Request failed: {}",
|
||||
e
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -99,12 +219,11 @@ impl Provider for AzureProvider {
|
||||
ProviderMetadata::new(
|
||||
"azure_openai",
|
||||
"Azure OpenAI",
|
||||
"Models through Azure OpenAI Service",
|
||||
"Models through Azure OpenAI Service (uses Azure credential chain by default)",
|
||||
"gpt-4o",
|
||||
AZURE_OPENAI_KNOWN_MODELS.to_vec(),
|
||||
AZURE_DOC_URL,
|
||||
vec![
|
||||
ConfigKey::new("AZURE_OPENAI_API_KEY", true, true, None),
|
||||
ConfigKey::new("AZURE_OPENAI_ENDPOINT", true, false, None),
|
||||
ConfigKey::new("AZURE_OPENAI_DEPLOYMENT_NAME", true, false, None),
|
||||
ConfigKey::new("AZURE_OPENAI_API_VERSION", true, false, Some("2024-10-21")),
|
||||
|
||||
Reference in New Issue
Block a user