mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: address OCR bridge review comments
This commit is contained in:
parent
e55b9f981a
commit
82ec9aeb70
8 changed files with 151 additions and 49 deletions
|
|
@ -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}")]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue