fix: address OCR bridge review comments

This commit is contained in:
Ishaan Jaff 2026-06-24 16:04:31 -07:00
parent e55b9f981a
commit 82ec9aeb70
No known key found for this signature in database
8 changed files with 151 additions and 49 deletions

View file

@ -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}")]

View file

@ -53,11 +53,21 @@ fn ocr_config_for(provider: LlmProvider) -> Option<&'static dyn OcrProviderConfi
}
}
fn string_headers(extra_headers: Option<Map<String, Value>>) -> Vec<(String, String)> {
fn string_headers(extra_headers: Option<Map<String, Value>>) -> CoreResult<Vec<(String, String)>> {
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<Value> {
.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()
)
);
}
}

View file

@ -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()),
}

View file

@ -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,

View file

@ -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(

View file

@ -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,

View file

@ -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:

View file

@ -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 = []