fix(rust): honor vertex_project, vertex_location and enable_azure_ad_token_refresh globals

Python resolves the Vertex project and location as call params, then the
litellm.vertex_project / litellm.vertex_location globals, then env, and
Azure AD token refresh from litellm.enable_azure_ad_token_refresh alone.
Native OCR skipped the globals, so a config.yaml litellm_settings value
silently fell through to the credential's project and us-central1, and a
managed identity setup without an API key failed. The bridge now reads
them through a provider_defaults settings group into OcrSettings, and
VertexConfig / AzureAuthInputs slot them in at Python's precedence.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Yujong Lee 2026-09-18 21:01:37 -07:00
parent 0d76359dc9
commit 1ee4b62e9c
17 changed files with 221 additions and 51 deletions

View file

@ -1976,6 +1976,7 @@ dependencies = [
"azure_identity",
"litellm-auth",
"moka",
"rstest",
"serde_json",
"sha2 0.10.9",
"strum",

View file

@ -18,4 +18,5 @@ azure_core = "1.0.0"
azure_identity = { version = "1.0.0", features = ["tokio"] }
[dev-dependencies]
rstest.workspace = true
tokio.workspace = true

View file

@ -1,11 +1,10 @@
use serde_json::{Map, Value};
use std::collections::BTreeMap;
use strum::EnumString;
use litellm_auth::Error;
use litellm_auth::{
CredentialResolverHandle, InputSource, SecretValue, Sourced, TokenProviderHandle,
CredentialResolverHandle, Error, InputSource, SecretValue, Sourced, TokenProviderHandle,
};
use serde_json::{Map, Value};
use strum::EnumString;
pub const DEFAULT_AZURE_SCOPE: &str = "https://cognitiveservices.azure.com/.default";
@ -52,6 +51,16 @@ pub struct AzureAuthInputs {
}
impl AzureAuthInputs {
pub fn or_configured_token_refresh(self, enabled: bool) -> Self {
if *self.enable_azure_ad_token_refresh.value() || !enabled {
return self;
}
Self {
enable_azure_ad_token_refresh: Sourced::new(true, InputSource::Deployment),
..self
}
}
#[cfg(test)]
pub fn from_optional_params(params: &Map<String, Value>) -> Result<Self, Error> {
Self::from_sourced_optional_params(params, &BTreeMap::new())
@ -115,12 +124,12 @@ fn source_for(sources: &BTreeMap<String, InputSource>, name: &str) -> InputSourc
#[cfg(test)]
mod tests {
use serde_json::json;
use std::collections::BTreeMap;
use super::{AzureAuthInputs, AzureCredentialType, ConfigValue};
use litellm_auth::{InputSource, Sourced};
use serde_json::json;
use super::{AzureAuthInputs, AzureCredentialType, ConfigValue};
#[test]
fn selector_parsing_is_exact() {
@ -189,4 +198,29 @@ mod tests {
assert!(!debug.contains("token-value"));
assert!(!debug.contains("secret-value"));
}
#[rstest::rstest]
#[case::global_turns_refresh_on(json!({}), true, true, InputSource::Deployment)]
#[case::global_overrides_a_call_false_like_python(json!({"enable_azure_ad_token_refresh": false}), true, true, InputSource::Deployment)]
#[case::call_true_survives_a_global_false(json!({"enable_azure_ad_token_refresh": true}), false, true, InputSource::Request)]
#[case::both_off(json!({}), false, false, InputSource::Request)]
fn token_refresh_follows_the_configured_global(
#[case] params: serde_json::Value,
#[case] global: bool,
#[case] enabled: bool,
#[case] source: InputSource,
) {
let sources = BTreeMap::from([(
"enable_azure_ad_token_refresh".to_string(),
InputSource::Request,
)]);
let inputs =
AzureAuthInputs::from_sourced_optional_params(params.as_object().unwrap(), &sources)
.unwrap()
.or_configured_token_refresh(global);
assert_eq!(
inputs.enable_azure_ad_token_refresh,
Sourced::new(enabled, source)
);
}
}

View file

@ -1,17 +1,13 @@
use std::collections::BTreeMap;
use std::future::Future;
use std::path::Path;
use std::pin::Pin;
use std::sync::Arc;
use std::{collections::BTreeMap, future::Future, path::Path, pin::Pin, sync::Arc};
use gcp_auth::{CustomServiceAccount, TokenProvider};
use litellm_auth::{
CredentialPlacement, Error, InputSource, SecretValue, Sourced, http::apply_credential,
};
use moka::future::Cache;
use serde_json::{Map, Value};
use sha2::{Digest, Sha256};
use litellm_auth::http::apply_credential;
use litellm_auth::{CredentialPlacement, Error, InputSource, SecretValue, Sourced};
const CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform";
const GOOGLE_OAUTH_TOKEN_ENDPOINT: &str = "https://oauth2.googleapis.com/token";
const GOOGLE_APPLICATION_CREDENTIALS_ENV: &str = "GOOGLE_APPLICATION_CREDENTIALS";
@ -45,6 +41,16 @@ impl VertexConfig {
})
}
pub fn or_configured(self, project_id: Option<&str>, location: Option<&str>) -> Self {
let configured =
|value: Option<&str>| value.filter(|value| !value.is_empty()).map(str::to_string);
Self {
project_id: self.project_id.or_else(|| configured(project_id)),
location: self.location.or_else(|| configured(location)),
..self
}
}
pub fn project_id(&self) -> Option<&str> {
self.project_id.as_deref()
}
@ -571,4 +577,29 @@ mod tests {
assert_eq!(loads.load(Ordering::SeqCst), 1);
assert_eq!(calls.load(Ordering::SeqCst), 4);
}
#[test]
fn configured_defaults_sit_between_call_params_and_the_environment() {
let env = |name: &str| Some(format!("env-{name}"));
let from_config =
VertexConfig::default().or_configured(Some("global-project"), Some("global-location"));
assert_eq!(
get_vertex_ai_project(&from_config, &env).as_deref(),
Some("global-project")
);
assert_eq!(
get_vertex_ai_location(&from_config, &env).as_deref(),
Some("global-location")
);
let from_call =
config(json!({"vertex_project":"call-project","vertex_location":"call-location"}))
.or_configured(Some("global-project"), Some("global-location"));
assert_eq!(from_call.project_id(), Some("call-project"));
assert_eq!(from_call.location(), Some("call-location"));
let empty_global = VertexConfig::default().or_configured(Some(""), None);
assert_eq!(
get_vertex_ai_project(&empty_global, &env).as_deref(),
Some("env-VERTEXAI_PROJECT")
);
}
}

View file

@ -1,8 +1,8 @@
use litellm_auth::InputSource;
use litellm_llms::base_llm::ocr::transformation::OcrResponseFormat;
use litellm_llms::base_llm::ocr::{settings::OcrSettings, transformation::OcrResponseFormat};
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use super::test_support::{MockResponse, mock_server, ocr_client, perform_ocr, wire_request};
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
@ -48,6 +48,27 @@ async fn facade_executes_vertex_mistral_with_resolved_project_and_location() {
);
}
#[tokio::test]
async fn configured_project_and_location_apply_when_the_call_sets_neither() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let client = ocr_client().with_settings(OcrSettings {
vertex_project: Some("configured-project".into()),
vertex_location: Some("europe-west4".into()),
..OcrSettings::default()
});
crate::ocr::client::perform(
&client,
wire_request("vertex_ai/mistral-ocr-maas", &base, json!({})),
)
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].starts_with(
"POST /v1/projects/configured-project/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict "
));
}
#[tokio::test]
async fn supplied_authorization_is_forwarded_without_a_static_token() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;

View file

@ -3,7 +3,21 @@ use std::sync::OnceLock;
use litellm_auth::{InputSource, Sourced};
use litellm_auth_azure::{AzureAuthInputs, AzureAuthService};
use crate::base_llm::ocr::{error::Error, transformation::OcrConnection};
use crate::base_llm::ocr::{
error::Error,
transformation::{OcrConnection, PreparedOcrRequest},
};
pub(crate) fn azure_auth_inputs(request: &PreparedOcrRequest) -> Result<AzureAuthInputs, Error> {
Ok(AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?
}
.or_configured_token_refresh(request.connection.settings.enable_azure_ad_token_refresh))
}
pub(super) async fn resolve_entra(
config: &AzureAuthInputs,

View file

@ -174,13 +174,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOcrConfig {
request: &PreparedOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, Error> {
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?
};
let config = crate::azure_ai::ocr::common_utils::azure_auth_inputs(request)?;
self.resolve_headers(&request.connection, &config, &|name: &str| {
request.connection.secret(name)
})

View file

@ -50,13 +50,7 @@ impl BaseOcrConfig for AzureAiOcrConfig {
request: &PreparedOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, Error> {
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?
};
let config = crate::azure_ai::ocr::common_utils::azure_auth_inputs(request)?;
self.resolve_headers(&request.connection, &config, &|name: &str| {
request.connection.secret(name)
})

View file

@ -11,6 +11,9 @@ pub struct OcrSettings {
pub poll_timeout: Duration,
pub document_intelligence_api_version: String,
pub document_intelligence_dpi: i64,
pub vertex_project: Option<String>,
pub vertex_location: Option<String>,
pub enable_azure_ad_token_refresh: bool,
}
impl Default for OcrSettings {
@ -21,6 +24,9 @@ impl Default for OcrSettings {
poll_timeout: Duration::from_secs(120),
document_intelligence_api_version: "2024-11-30".into(),
document_intelligence_dpi: 96,
vertex_project: None,
vertex_location: None,
enable_azure_ad_token_refresh: false,
}
}
}
@ -48,6 +54,7 @@ impl OcrSettings {
document_intelligence_dpi: env
.parsed("AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI")
.unwrap_or(defaults.document_intelligence_dpi),
..defaults
}
}
}
@ -96,6 +103,7 @@ mod tests {
poll_timeout: Duration::from_secs(600),
document_intelligence_api_version: "2025-01-01".into(),
document_intelligence_dpi: 72,
..OcrSettings::default()
}
);
}

View file

@ -1,6 +1,22 @@
use litellm_auth::InputSource;
use litellm_auth_gcp::VertexConfig;
use crate::base_llm::ocr::{error::Error, transformation::OcrConnection};
use crate::base_llm::ocr::{
error::Error,
transformation::{OcrConnection, PreparedOcrRequest},
};
pub(super) fn vertex_config(request: &PreparedOcrRequest) -> Result<VertexConfig, Error> {
let settings = &request.connection.settings;
Ok(VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?
.or_configured(
settings.vertex_project.as_deref(),
settings.vertex_location.as_deref(),
))
}
pub(super) fn validate_destination(connection: &OcrConnection) -> Result<(), Error> {
if connection.api_base.is_some() && connection.api_base_source == InputSource::Request {

View file

@ -1,9 +1,9 @@
use litellm_auth_gcp::{self as vertex, VertexConfig};
use litellm_auth_gcp as vertex;
use litellm_core_utils::{call_arguments::CallArguments, params::OpaqueParams, url_utils::ApiUrl};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use super::transformation::VertexAiOcrConfig;
use super::{common_utils::vertex_config, transformation::VertexAiOcrConfig};
use crate::base_llm::ocr::{
error::Error,
handler::OcrClient,
@ -122,10 +122,7 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
_params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, Error> {
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?;
let config = vertex_config(request)?;
let location =
vertex::get_vertex_ai_location(&config, &|name: &str| request.connection.secret(name))
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());

View file

@ -2,7 +2,7 @@ use litellm_auth_gcp::{self as vertex, VertexConfig};
use litellm_core_utils::{call_arguments::CallArguments, params::OpaqueParams, url_utils::ApiUrl};
use serde_json::Value;
use super::common_utils::validate_destination;
use super::common_utils::{validate_destination, vertex_config};
use crate::{
base_llm::ocr::{
document::{inline_remote_document, validate_inline_document},
@ -47,10 +47,7 @@ impl BaseOcrConfig for VertexAiOcrConfig {
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<Self::Environment, Error> {
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?;
let config = vertex_config(request)?;
self.resolve_environment(&request.connection, &config, client)
.await
}
@ -61,10 +58,7 @@ impl BaseOcrConfig for VertexAiOcrConfig {
_optional_params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, Error> {
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?;
let config = vertex_config(request)?;
let location =
vertex::get_vertex_ai_location(&config, &|name: &str| request.connection.secret(name))
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());

View file

@ -14,5 +14,10 @@
"url_policy": [
"user_url_validation",
"user_url_allowed_hosts"
],
"provider_defaults": [
"vertex_project",
"vertex_location",
"enable_azure_ad_token_refresh"
]
}

View file

@ -7,16 +7,18 @@ const MODULE: &str = "litellm.rust_bridge.settings";
pub(crate) enum PythonSettings {
Http,
UrlPolicy,
ProviderDefaults,
}
impl PythonSettings {
#[cfg(test)]
pub(crate) const ALL: [Self; 2] = [Self::Http, Self::UrlPolicy];
pub(crate) const ALL: [Self; 3] = [Self::Http, Self::UrlPolicy, Self::ProviderDefaults];
pub(crate) fn name(self) -> &'static str {
match self {
Self::Http => "http_settings",
Self::UrlPolicy => "url_policy",
Self::ProviderDefaults => "provider_defaults",
}
}

View file

@ -16,7 +16,11 @@ use pyo3::{
types::{PyDict, PyTuple},
};
use crate::{errors::RustBridgeDeclined, http, python_settings::PythonSecrets};
use crate::{
errors::RustBridgeDeclined,
http,
python_settings::{PythonSecrets, PythonSettings},
};
const SURFACE: LegacySurface = LegacySurface {
call_type: "ocr",
@ -44,7 +48,7 @@ fn run_ocr(
&config,
http::url_policy(py)?,
VERTEX_AUTH.clone(),
OcrSettings::from_environment(&ProcessEnvironment),
ocr_settings(py)?,
Arc::new(PythonSecrets),
)
.map_err(|error| RustBridgeDeclined::new_err(error.to_string()))?;
@ -58,6 +62,30 @@ fn run_ocr(
)
}
#[derive(FromPyObject)]
struct PythonProviderDefaults {
vertex_project: Option<String>,
vertex_location: Option<String>,
enable_azure_ad_token_refresh: Option<bool>,
}
fn ocr_settings(py: Python<'_>) -> PyResult<OcrSettings> {
let defaults: PythonProviderDefaults = PythonSettings::ProviderDefaults
.read(py)?
.extract()
.map_err(|error: PyErr| {
RustBridgeDeclined::new_err(format!(
"litellm provider defaults cannot be used by the Rust route: {error}"
))
})?;
Ok(OcrSettings {
vertex_project: defaults.vertex_project,
vertex_location: defaults.vertex_location,
enable_azure_ad_token_refresh: defaults.enable_azure_ad_token_refresh == Some(true),
..OcrSettings::from_environment(&ProcessEnvironment)
})
}
#[pyfunction]
pub(crate) fn ocr(
py: Python<'_>,

View file

@ -24,6 +24,13 @@ class UrlPolicy:
user_url_allowed_hosts: Sequence[str]
@dataclass(frozen=True, slots=True)
class ProviderDefaults:
vertex_project: str | None
vertex_location: str | None
enable_azure_ad_token_refresh: bool | None
def warn(message: str) -> None:
from litellm._logging import verbose_logger
@ -36,6 +43,16 @@ def secret(name: str) -> str | None:
return get_secret_str(name)
def provider_defaults() -> ProviderDefaults:
import litellm
return ProviderDefaults(
vertex_project=litellm.vertex_project,
vertex_location=litellm.vertex_location,
enable_azure_ad_token_refresh=litellm.enable_azure_ad_token_refresh,
)
def url_policy() -> UrlPolicy:
import litellm

View file

@ -22,6 +22,7 @@ def test_the_rust_contract_matches_the_returned_fields() -> None:
assert contract == {
"http_settings": [field.name for field in dataclasses.fields(settings.http_settings())],
"url_policy": [field.name for field in dataclasses.fields(settings.url_policy())],
"provider_defaults": [field.name for field in dataclasses.fields(settings.provider_defaults())],
}
@ -119,3 +120,15 @@ def test_secret_reads_the_environment_without_a_secret_manager(monkeypatch: pyte
monkeypatch.setattr(litellm, "secret_manager_client", None)
assert settings.secret("MISTRAL_API_KEY") == "env-key"
def test_provider_defaults_read_the_litellm_globals(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "vertex_project", "configured-project")
monkeypatch.setattr(litellm, "vertex_location", "europe-west4")
monkeypatch.setattr(litellm, "enable_azure_ad_token_refresh", True)
assert settings.provider_defaults() == settings.ProviderDefaults(
vertex_project="configured-project",
vertex_location="europe-west4",
enable_azure_ad_token_refresh=True,
)