diff --git a/litellm/llms/base_llm/ocr/transformation.py b/litellm/llms/base_llm/ocr/transformation.py index 3668c64defa..280a8e43049 100644 --- a/litellm/llms/base_llm/ocr/transformation.py +++ b/litellm/llms/base_llm/ocr/transformation.py @@ -78,7 +78,7 @@ class OCRResponse(LiteLLMPydanticObjectBase): _hidden_params: dict = PrivateAttr(default_factory=dict) def model_post_init(self, __context: Any) -> None: - extra_fields = getattr(self, "__pydantic_extra__", None) + extra_fields = self.model_extra if isinstance(extra_fields, dict): hidden_params = extra_fields.pop("_hidden_params", None) if isinstance(hidden_params, dict): diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 6cd8b641b34..9ddd713d5fa 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -266,6 +266,8 @@ def _get_rust_bridge_callbacks() -> tuple[list[object], list[object]]: if isinstance(callback, CustomGuardrail): if callback not in guardrails: guardrails.append(callback) + # CustomGuardrail inherits CustomLogger; keep it in both lists so + # Rust OCR runs guard hooks and success/failure logging hooks. if isinstance(callback, CustomLogger) and callback not in callbacks: callbacks.append(callback) return callbacks, guardrails diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 48493fec7ab..3be8482f06b 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -390,7 +390,9 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs ), litellm_overhead_time_ms=litellm_overhead_time_ms, litellm_rust=( - standard_logging_payload.get("hidden_params", {}).get("litellm_rust", None) + (standard_logging_payload.get("hidden_params") or {}).get( + "litellm_rust", None + ) if standard_logging_payload is not None else None ), diff --git a/tests/ocr_tests/test_ocr_rust_bridge_callbacks.py b/tests/ocr_tests/test_ocr_rust_bridge_callbacks.py index 6393f6381b5..0fb7dab7eda 100644 --- a/tests/ocr_tests/test_ocr_rust_bridge_callbacks.py +++ b/tests/ocr_tests/test_ocr_rust_bridge_callbacks.py @@ -1,4 +1,4 @@ -"""Tests proving Rust OCR receives and executes Python callbacks.""" +"""Network-free tests proving Rust OCR receives and executes Python callbacks.""" from datetime import datetime from typing import Any, Optional @@ -14,7 +14,7 @@ from litellm.types.utils import StandardLoggingPayload MODEL = "mistral/mistral-ocr-latest" DOCUMENT: dict[str, object] = { "type": "document_url", - "document_url": "https://example.com/test.pdf", + "document_url": "data:application/pdf;base64,JVBERi0xLjQK", } FAKE_RUST_OCR_RESPONSE: dict[str, object] = { "pages": [{"index": 0, "markdown": "hello from rust ocr"}], @@ -78,7 +78,7 @@ class OCRCustomGuardrail(CustomGuardrail): self.success_log_calls += 1 -class ExecutingRustAocrBridge: +class MockExecutingRustAocrBridge: def __init__(self) -> None: self.calls: list[dict[str, object]] = [] @@ -155,7 +155,7 @@ def reset_litellm_rust_callbacks(): @pytest.mark.asyncio async def test_rust_ocr_executes_custom_logger_from_callback_manager(): - bridge = ExecutingRustAocrBridge() + bridge = MockExecutingRustAocrBridge() custom_logger = OCRCustomLogger() litellm.logging_callback_manager.add_litellm_callback(custom_logger) litellm.use_litellm_rust(True, aocr=bridge) @@ -188,7 +188,7 @@ async def test_rust_ocr_executes_custom_logger_from_callback_manager(): @pytest.mark.asyncio async def test_rust_ocr_executes_custom_guardrail_from_callback_manager(): - bridge = ExecutingRustAocrBridge() + bridge = MockExecutingRustAocrBridge() custom_guardrail = OCRCustomGuardrail() litellm.logging_callback_manager.add_litellm_callback(custom_guardrail) litellm.use_litellm_rust(True, aocr=bridge)