diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index b57fb2f5de8..c05bd6103e9 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -13,6 +13,8 @@ pub enum CoreError { MissingField(&'static str), #[error("invalid response: {0}")] InvalidResponse(String), + #[error("invalid request: {0}")] + InvalidRequest(String), #[error("invalid provider: {0}")] InvalidProvider(String), #[error("{0}")] diff --git a/litellm-rust/crates/providers/src/ocr.rs b/litellm-rust/crates/providers/src/ocr.rs index a52d5532afc..17d5dc1ac83 100644 --- a/litellm-rust/crates/providers/src/ocr.rs +++ b/litellm-rust/crates/providers/src/ocr.rs @@ -53,11 +53,21 @@ fn ocr_config_for(provider: LlmProvider) -> Option<&'static dyn OcrProviderConfi } } -fn string_headers(extra_headers: Option>) -> Vec<(String, String)> { +fn string_headers(extra_headers: Option>) -> CoreResult> { extra_headers .unwrap_or_default() .into_iter() - .filter_map(|(key, value)| value.as_str().map(|value| (key, value.to_string()))) + .map(|(key, value)| { + value + .as_str() + .map(|value| (key.clone(), value.to_string())) + .ok_or_else(|| { + CoreError::InvalidRequest(format!( + "OCR extra_headers.{key} must be a string, got {}", + litellm_core::error::json_type_name(&value) + )) + }) + }) .collect() } @@ -93,7 +103,7 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult { .data; let mut request_builder = http_client().post(&url).bearer_auth(&api_key).json(&body); - for (key, value) in string_headers(request.extra_headers) { + for (key, value) in string_headers(request.extra_headers)? { request_builder = request_builder.header(&key, value); } if let Some(duration) = request.timeout { @@ -165,19 +175,35 @@ mod tests { } #[test] - fn string_headers_keeps_only_string_values() { + fn string_headers_accepts_string_values() { let headers = json!({ - "x-trace-id": "trace-1", - "x-number": 42, - "x-bool": true + "x-trace-id": "trace-1" }) .as_object() .unwrap() .clone(); assert_eq!( - string_headers(Some(headers)), + string_headers(Some(headers)).expect("string headers accepted"), vec![("x-trace-id".to_string(), "trace-1".to_string())] ); } + + #[test] + fn string_headers_rejects_non_string_values() { + let headers = json!({ + "x-retry-count": 3 + }) + .as_object() + .unwrap() + .clone(); + + let err = string_headers(Some(headers)).expect_err("non-string header rejected"); + assert_eq!( + err, + CoreError::InvalidRequest( + "OCR extra_headers.x-retry-count must be a string, got number".to_string() + ) + ); + } } diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 46ff700f456..aa1fb2e0127 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -36,6 +36,7 @@ fn core_error_to_pyerr(err: CoreError) -> PyErr { CoreError::Auth(message) => PyValueError::new_err(message), CoreError::InvalidProvider(_) | CoreError::InvalidType { .. } + | CoreError::InvalidRequest(_) | CoreError::MissingField(_) => PyValueError::new_err(err.to_string()), other => PyRuntimeError::new_err(other.to_string()), } diff --git a/litellm/llms/base_llm/ocr/transformation.py b/litellm/llms/base_llm/ocr/transformation.py index 263e0c094ce..de6bc2471ec 100644 --- a/litellm/llms/base_llm/ocr/transformation.py +++ b/litellm/llms/base_llm/ocr/transformation.py @@ -101,6 +101,12 @@ class BaseOCRConfig: """ return [] + def get_api_key_env_var(self) -> Optional[str]: + """ + Return the provider-specific API key environment variable name, if any. + """ + return None + def map_ocr_params( self, non_default_params: dict, diff --git a/litellm/llms/mistral/ocr/transformation.py b/litellm/llms/mistral/ocr/transformation.py index 21e0e27a314..49555baaf70 100644 --- a/litellm/llms/mistral/ocr/transformation.py +++ b/litellm/llms/mistral/ocr/transformation.py @@ -15,6 +15,8 @@ from litellm.llms.base_llm.ocr.transformation import ( ) from litellm.secret_managers.main import get_secret_str +MISTRAL_OCR_API_KEY_ENV_VAR = "MISTRAL_API_KEY" + class MistralOCRConfig(BaseOCRConfig): """ @@ -59,6 +61,9 @@ class MistralOCRConfig(BaseOCRConfig): "id", ] + def get_api_key_env_var(self) -> Optional[str]: + return MISTRAL_OCR_API_KEY_ENV_VAR + def map_ocr_params( self, non_default_params: dict, @@ -95,7 +100,7 @@ class MistralOCRConfig(BaseOCRConfig): """ # Get API key from environment if not provided if api_key is None: - api_key = get_secret_str("MISTRAL_API_KEY") + api_key = get_secret_str(MISTRAL_OCR_API_KEY_ENV_VAR) if api_key is None: raise ValueError( diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 98f6c6af8f2..f79c17256ca 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -49,6 +49,13 @@ class _PreparedOCRRequest: litellm_logging_obj: LiteLLMLoggingObj +@dataclass +class _PreparedRustOCRCall: + api_key: Optional[str] + headers: dict[str, object] + complete_url: str + + def _timeout_to_seconds( timeout: Optional[Union[float, httpx.Timeout]], ) -> Optional[float]: @@ -166,8 +173,7 @@ def _prepare_ocr_request( ) -def _run_rust_ocr( - rust_ocr: RustOcr, +def _prepare_rust_ocr_call( logging_obj: LiteLLMLoggingObj, provider_config: BaseOCRConfig, resolve_api_key: Callable[[str], Optional[str]], @@ -175,21 +181,14 @@ def _run_rust_ocr( document: dict[str, object], api_key: Optional[str], api_base: Optional[str], - custom_llm_provider: str, extra_headers: Optional[dict[str, object]], optional_params: dict[str, object], litellm_params: dict[str, object], - timeout_seconds: Optional[float], -) -> OCRResponse: - """Run the Mistral OCR call through the Rust bridge and wrap the result. - - Resolves the key the same way the Python path does so secret-manager backends - (AWS/Azure/GCP/Vault) work; the Rust bridge's own fallback only reads the - process environment. The request that Rust actually sends (resolved URL and - headers) is mirrored into pre_call so logs match the wire. Dependencies are - injected so this stays unit-testable without patching module globals. - """ - resolved_api_key = api_key or resolve_api_key("MISTRAL_API_KEY") +) -> _PreparedRustOCRCall: + api_key_env_var = provider_config.get_api_key_env_var() + resolved_api_key = api_key or ( + resolve_api_key(api_key_env_var) if api_key_env_var is not None else None + ) resolved_headers = provider_config.validate_environment( headers=extra_headers or {}, model=model, @@ -216,11 +215,53 @@ def _run_rust_ocr( "headers": resolved_headers, }, ) + return _PreparedRustOCRCall( + api_key=resolved_api_key, + headers=cast(dict[str, object], resolved_headers), + complete_url=resolved_complete_url, + ) + + +def _run_rust_ocr( + rust_ocr: RustOcr, + logging_obj: LiteLLMLoggingObj, + provider_config: BaseOCRConfig, + resolve_api_key: Callable[[str], Optional[str]], + model: str, + document: dict[str, object], + api_key: Optional[str], + api_base: Optional[str], + custom_llm_provider: str, + extra_headers: Optional[dict[str, object]], + optional_params: dict[str, object], + litellm_params: dict[str, object], + timeout_seconds: Optional[float], +) -> OCRResponse: + """Run the Mistral OCR call through the Rust bridge and wrap the result. + + Resolves the key the same way the Python path does so secret-manager backends + (AWS/Azure/GCP/Vault) work; the Rust bridge's own fallback only reads the + process environment. The request that Rust actually sends (resolved URL and + headers) is mirrored into pre_call so logs match the wire. Dependencies are + injected so this stays unit-testable without patching module globals. + """ + prepared = _prepare_rust_ocr_call( + logging_obj=logging_obj, + provider_config=provider_config, + resolve_api_key=resolve_api_key, + model=model, + document=document, + api_key=api_key, + api_base=api_base, + extra_headers=extra_headers, + optional_params=optional_params, + litellm_params=litellm_params, + ) return OCRResponse.model_validate( rust_ocr( model=model, document=document, - api_key=resolved_api_key, + api_key=prepared.api_key, api_base=api_base, custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, @@ -245,38 +286,23 @@ async def _run_rust_aocr( litellm_params: dict[str, object], timeout_seconds: Optional[float], ) -> OCRResponse: - resolved_api_key = api_key or resolve_api_key("MISTRAL_API_KEY") - resolved_headers = provider_config.validate_environment( - headers=extra_headers or {}, + prepared = _prepare_rust_ocr_call( + logging_obj=logging_obj, + provider_config=provider_config, + resolve_api_key=resolve_api_key, model=model, - api_key=resolved_api_key, + document=document, + api_key=api_key, api_base=api_base, - litellm_params=litellm_params, - ) - resolved_complete_url = provider_config.get_complete_url( - api_base=api_base, - model=model, + extra_headers=extra_headers, optional_params=optional_params, litellm_params=litellm_params, ) - logging_obj.pre_call( - input="OCR document processing", - api_key=resolved_api_key, - additional_args={ - "complete_input_dict": { - "model": model, - "document": document, - **optional_params, - }, - "api_base": resolved_complete_url, - "headers": resolved_headers, - }, - ) return OCRResponse.model_validate( await rust_aocr( model=model, document=document, - api_key=resolved_api_key, + api_key=prepared.api_key, api_base=api_base, custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, diff --git a/litellm/ocr/rust_bridge.py b/litellm/ocr/rust_bridge.py index 0dec57d9168..1e3312c1473 100644 --- a/litellm/ocr/rust_bridge.py +++ b/litellm/ocr/rust_bridge.py @@ -27,7 +27,8 @@ class RustOcr(Protocol): extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout_seconds: float | None, - ) -> dict[str, object]: ... + ) -> dict[str, object]: + raise NotImplementedError class RustAocr(Protocol): @@ -43,7 +44,8 @@ class RustAocr(Protocol): extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: ... + ) -> Awaitable[dict[str, object]]: + raise NotImplementedError class _Unset: diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index aee8506b84b..d51b56330d2 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -109,6 +109,12 @@ class RecordingLogging: class FakeOCRConfig: """A stand-in ``BaseOCRConfig`` that echoes the request it would build.""" + def __init__(self, api_key_env_var="MISTRAL_API_KEY"): + self.api_key_env_var = api_key_env_var + + def get_api_key_env_var(self): + return self.api_key_env_var + def validate_environment( self, *, headers, model, api_key, api_base, litellm_params ): @@ -280,6 +286,34 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): assert bridge.calls[0]["api_key"] == "sk-from-vault" +def test_run_rust_ocr_uses_provider_api_key_env_var(): + bridge = RecordingBridge() + resolver_calls = [] + + def _resolver(name): + resolver_calls.append(name) + return "sk-provider-env" + + ocr_main._run_rust_ocr( + rust_ocr=bridge, + logging_obj=RecordingLogging(), + provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"), + resolve_api_key=_resolver, + model="provider-ocr-model", + document=DOCUMENT, + api_key=None, + api_base=None, + custom_llm_provider="mistral", + extra_headers=None, + optional_params={}, + litellm_params={}, + timeout_seconds=None, + ) + + assert resolver_calls == ["PROVIDER_OCR_API_KEY"] + assert bridge.calls[0]["api_key"] == "sk-provider-env" + + def test_run_rust_ocr_prefers_explicit_key_over_resolver(): bridge = RecordingBridge() resolver_calls = []