From 4da89566eb1e9e55cd7e1a558095aa03ce526cb2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 16 Jul 2026 21:05:52 -0700 Subject: [PATCH] feat(ocr): resolve environment references in Rust (#33598) Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm-rust/crates/ai-gateway/src/config.rs | 85 ++++++++++++++++ .../crates/ai-gateway/src/constants.rs | 8 ++ litellm-rust/crates/ai-gateway/src/io/ocr.rs | 19 +++- .../crates/ai-gateway/src/io/ocr/tests.rs | 62 ++++++++++++ litellm-rust/crates/ai-gateway/src/lib.rs | 5 +- tests/e2e/gateway/test_ocr_rust_e2e.py | 96 ++++++++++++++++++- tests/test_litellm/ocr/test_rust_bridge.py | 47 +++++++++ 7 files changed, 315 insertions(+), 7 deletions(-) create mode 100644 litellm-rust/crates/ai-gateway/src/config.rs diff --git a/litellm-rust/crates/ai-gateway/src/config.rs b/litellm-rust/crates/ai-gateway/src/config.rs new file mode 100644 index 00000000000..a98af8fb794 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/config.rs @@ -0,0 +1,85 @@ +use crate::constants::ENV_REFERENCE_PREFIX; + +pub(crate) fn resolve_env_reference( + value: Option<&str>, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Option { + let value = value?; + let Some(name) = value.strip_prefix(ENV_REFERENCE_PREFIX) else { + return Some(value.to_string()); + }; + if name.trim().is_empty() { + return None; + } + env_lookup(name).filter(|resolved| !resolved.trim().is_empty()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn env_lookup(name: &str) -> Option { + match name { + "PRESENT" => Some("resolved".to_string()), + "BLANK" => Some(" ".to_string()), + _ => None, + } + } + + #[test] + fn preserves_explicit_value() { + assert_eq!( + resolve_env_reference(Some("explicit"), &env_lookup), + Some("explicit".to_string()) + ); + } + + #[test] + fn preserves_value_that_only_contains_reference_prefix() { + assert_eq!( + resolve_env_reference(Some("prefix-os.environ/PRESENT"), &env_lookup), + Some("prefix-os.environ/PRESENT".to_string()) + ); + } + + #[test] + fn resolves_present_reference() { + assert_eq!( + resolve_env_reference(Some("os.environ/PRESENT"), &env_lookup), + Some("resolved".to_string()) + ); + } + + #[test] + fn missing_reference_is_absent() { + assert_eq!( + resolve_env_reference(Some("os.environ/MISSING"), &env_lookup), + None + ); + } + + #[test] + fn blank_reference_value_is_absent() { + assert_eq!( + resolve_env_reference(Some("os.environ/BLANK"), &env_lookup), + None + ); + } + + #[test] + fn malformed_reference_is_absent() { + assert_eq!( + resolve_env_reference(Some("os.environ/"), &env_lookup), + None + ); + assert_eq!( + resolve_env_reference(Some("os.environ/ "), &env_lookup), + None + ); + } + + #[test] + fn absent_input_stays_absent() { + assert_eq!(resolve_env_reference(None, &env_lookup), None); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 3116a4c9932..10689943bff 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -7,23 +7,31 @@ /// Default LiteLLM control-plane base URL for request-log egress when /// `LITELLM_PROXY_BASE_URL` is unset. +#[cfg(feature = "server")] pub(crate) const DEFAULT_PROXY_BASE_URL: &str = "http://localhost:4000"; /// The logs ingest path appended to the proxy base. Not a tunable; it is the /// proxy's API contract (the rust-control-plane router on the Python proxy). +#[cfg(feature = "server")] pub(crate) const RUST_CONTROL_PLANE_LOGS_PATH: &str = "/v1/rust_control_plane/logs"; /// Default bounded channel depth for the log-egress worker. /// Override: `LITELLM_LOG_CHANNEL_CAPACITY`. +#[cfg(feature = "server")] pub(crate) const DEFAULT_CHANNEL_CAPACITY: usize = 4096; /// Default max records POSTed per request to the control plane. /// Override: `LITELLM_LOG_BATCH_SIZE`. +#[cfg(feature = "server")] pub(crate) const DEFAULT_MAX_BATCH_SIZE: usize = 256; /// Default partial-batch flush cadence, in ms. /// Override: `LITELLM_LOG_FLUSH_INTERVAL_MS`. +#[cfg(feature = "server")] pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500; /// Provider attributed to realtime sessions in the logging payload. +#[cfg(feature = "server")] pub(crate) const DEFAULT_PROVIDER: &str = "openai"; + +pub(crate) const ENV_REFERENCE_PREFIX: &str = "os.environ/"; diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr.rs b/litellm-rust/crates/ai-gateway/src/io/ocr.rs index d233babb748..0d42125bdbd 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr.rs @@ -14,6 +14,8 @@ use litellm_core::ocr::transformation::{ use litellm_core::CoreResult; use serde_json::{Map, Value}; +use crate::config::resolve_env_reference; + mod common_utils; use common_utils::{ @@ -73,6 +75,14 @@ pub struct OcrRequest<'a> { /// /// Async: intended to be awaited directly by the Python bridge's async entrypoint. pub async fn ocr(request: OcrRequest<'_>) -> CoreResult { + let env_lookup = |key: &str| std::env::var(key).ok(); + ocr_with_env(request, &env_lookup).await +} + +async fn ocr_with_env( + request: OcrRequest<'_>, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> CoreResult { let model = request.model; let config = ocr_provider_config(request.custom_llm_provider, model).ok_or_else(|| { CoreError::InvalidProvider(format!( @@ -80,18 +90,19 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult { request.custom_llm_provider )) })?; - let env_lookup = |key: &str| std::env::var(key).ok(); + let api_key = resolve_env_reference(request.api_key, env_lookup); + let api_base = resolve_env_reference(request.api_base, env_lookup); let headers = string_headers(request.extra_headers)?; let auth_strategy = config.auth_strategy(); let api_key = (!has_header(&headers, auth_strategy.header_name())) - .then(|| config.resolve_api_key(request.api_key, &env_lookup)) + .then(|| config.resolve_api_key(api_key.as_deref(), env_lookup)) .transpose()?; let url = config.complete_url( - request.api_base, + api_base.as_deref(), model, &request.optional_params, - &env_lookup, + env_lookup, )?; let filtered_params = config.map_ocr_params(&request.optional_params); let upstream_headers = upstream_headers(&headers, auth_strategy, api_key.as_deref()); diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs index e2e9c65ab35..901e20275bc 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs @@ -198,6 +198,68 @@ async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() { ); } +#[tokio::test] +async fn ocr_resolves_api_key_and_base_references_in_rust() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("test listener binds"); + let addr = listener.local_addr().expect("listener has local addr"); + let resolved_base = format!("http://{addr}"); + + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts one request"); + let request = read_http_headers(&mut socket).await; + let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#; + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + response_body.len(), + response_body + ); + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + request + }); + + let env_lookup = |name: &str| match name { + "OCR_TEST_API_KEY" => Some("sk-resolved".to_string()), + "OCR_TEST_API_BASE" => Some(resolved_base.clone()), + _ => None, + }; + + let response = ocr_with_env( + OcrRequest { + model: "mistral-ocr-latest", + document: json!({ + "type": "document_url", + "document_url": "https://example.com/doc.pdf" + }), + api_key: Some("os.environ/OCR_TEST_API_KEY"), + api_base: Some("os.environ/OCR_TEST_API_BASE"), + custom_llm_provider: "mistral", + extra_headers: None, + optional_params: Map::new(), + timeout: Some(Duration::from_secs(5)), + }, + &env_lookup, + ) + .await + .expect("ocr request succeeds"); + + assert_eq!(response["pages"][0]["markdown"], "ok"); + + let request = server.await.expect("server task completes"); + assert!(request.starts_with("POST /v1/ocr HTTP/1.1"), "{request}"); + assert!( + request + .to_ascii_lowercase() + .contains("authorization: bearer sk-resolved"), + "{request}" + ); + assert!(!request.contains("os.environ/"), "{request}"); +} + #[tokio::test] async fn ocr_forwards_full_mistral_contract_and_filters_internal_params() { let listener = TcpListener::bind("127.0.0.1:0") diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index 6c04fbb7626..636a6d9252b 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -13,6 +13,9 @@ pub mod io; +mod config; +mod constants; + /// GIL-activity tracking. Pure (atomics only); shared by the `server` routes and /// the `python-config` reader, so it is available without either feature. pub mod gil; @@ -28,8 +31,6 @@ pub mod state; // `server`-gated; `io::realtime` exposes the generic `observe` hook while the // collector and callback fan-out live here. #[cfg(feature = "server")] -mod constants; -#[cfg(feature = "server")] pub mod integrations; #[cfg(feature = "server")] mod realtime; diff --git a/tests/e2e/gateway/test_ocr_rust_e2e.py b/tests/e2e/gateway/test_ocr_rust_e2e.py index f5d77afd3be..6f799ed80fc 100644 --- a/tests/e2e/gateway/test_ocr_rust_e2e.py +++ b/tests/e2e/gateway/test_ocr_rust_e2e.py @@ -10,6 +10,8 @@ from __future__ import annotations import json import os +import time +import uuid from dataclasses import dataclass from pathlib import Path from typing import Literal, cast @@ -126,10 +128,17 @@ class OcrResponseEnvelope(BaseModel): usage_info: OcrUsageInfo | None = None +class ModelInfoDetail(BaseModel): + model_config = ConfigDict(frozen=True, extra="allow") + + id: str | None = None + + class ModelInfoEntry(BaseModel): model_config = ConfigDict(frozen=True, extra="allow") model_name: str + model_info: ModelInfoDetail = Field(default_factory=ModelInfoDetail) class ModelInfoResponse(BaseModel): @@ -149,7 +158,6 @@ class GatewayConfig(BaseModel): model_list: tuple[GatewayConfigEntry, ...] - TEST_PDF_URL = ( "https://cdn.jsdelivr.net/gh/BerriAI/litellm" "@d769e81c90d453240c61fc572cdb27fae06a89d0" @@ -322,6 +330,55 @@ class OcrGateway: content=_wire_payload(model, document, params), ) + def create_model( + self, model_name: str, litellm_params: dict[str, str] + ) -> httpx.Response: + with self._client() as client: + return client.post( + f"{self.base_url.rstrip('/')}/model/new", + headers={"Authorization": f"Bearer {self.master_key}"}, + json={"model_name": model_name, "litellm_params": litellm_params}, + ) + + def delete_model(self, model_id: str) -> httpx.Response: + with self._client() as client: + return client.post( + f"{self.base_url.rstrip('/')}/model/delete", + headers={"Authorization": f"Bearer {self.master_key}"}, + json={"id": model_id}, + ) + + def model_id(self, model_name: str) -> str | None: + with self._client() as client: + response = client.get( + f"{self.base_url.rstrip('/')}/model/info", + headers={"Authorization": f"Bearer {self.master_key}"}, + ) + assert response.status_code == 200, response.text + parsed = ModelInfoResponse.model_validate_json(response.content) + for entry in parsed.data: + if entry.model_name == model_name: + return entry.model_info.id + return None + + def wait_for_model(self, model_name: str, attempts: int = 20) -> None: + for _ in range(attempts): + if model_name in self.model_names(): + return + time.sleep(1) + raise AssertionError( + f"{model_name} did not appear on /model/info within {attempts}s" + ) + + def wait_for_model_absent(self, model_name: str, attempts: int = 20) -> None: + for _ in range(attempts): + if model_name not in self.model_names(): + return + time.sleep(1) + raise AssertionError( + f"{model_name} still present on /model/info after {attempts}s" + ) + @dataclass(frozen=True) class OcrResources: @@ -410,3 +467,40 @@ def test_rust_ocr_proxy_forwards_full_contract_to_capture_endpoint( capture_proxy.captures.get(timeout=10) ) assert captured == EXPECTED_UPSTREAM + + +class TestRustOcrDynamicDeployment: + def test_os_environ_api_key_deployment_lifecycle( + self, resources: OcrResources + ) -> None: + if not os.getenv("MISTRAL_API_KEY"): + pytest.skip("Set MISTRAL_API_KEY on the proxy for the live OCR lifecycle") + + gateway = resources.gateway + model_name = f"rust-ocr-env-e2e-{uuid.uuid4().hex[:8]}" + + create = gateway.create_model( + model_name=model_name, + litellm_params={ + "model": "mistral/mistral-ocr-latest", + "api_key": "os.environ/MISTRAL_API_KEY", + }, + ) + assert create.status_code == 200, create.text + + try: + gateway.wait_for_model(model_name) + + response = gateway.ocr( + model_name, + OcrDocument(type="document_url", document_url=TEST_PDF_URL), + ) + assert response.status_code == 200, response.text + OcrResponseEnvelope.model_validate_json(response.content) + assert "os.environ/MISTRAL_API_KEY" not in response.text + finally: + deployed_id = gateway.model_id(model_name) + if deployed_id is not None: + delete = gateway.delete_model(deployed_id) + assert delete.status_code == 200, delete.text + gateway.wait_for_model_absent(model_name) diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index c0897a642a6..8288b0196ac 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -859,3 +859,50 @@ def test_raise_ocr_exception_keeps_validation_error_off_bad_request( ) assert spy.calls[0].original_exception is validation_info.value + + +def test_ocr_forwards_os_environ_api_key_reference_to_rust( + fake_bridge: RecordingBridge, +) -> None: + litellm.ocr( + model=MODEL, document=DOCUMENT, api_key="os.environ/MISTRAL_OCR_TEST_KEY" + ) + + assert fake_bridge.calls[0]["api_key"] == "os.environ/MISTRAL_OCR_TEST_KEY" + + +def test_ocr_forwards_provider_derived_os_environ_references_to_rust( + fake_bridge: RecordingBridge, monkeypatch: pytest.MonkeyPatch +) -> None: + def fake_get_llm_provider( + *, + model: str, + custom_llm_provider: str | None, + api_base: str | None, + api_key: str | None, + ) -> tuple[str, str, str, str]: + return ( + "mistral-ocr-latest", + "mistral", + "os.environ/MISTRAL_PROVIDER_KEY", + "os.environ/MISTRAL_PROVIDER_BASE", + ) + + monkeypatch.setattr(ocr_main.litellm, "get_llm_provider", fake_get_llm_provider) + + litellm.ocr(model=MODEL, document=DOCUMENT) + + call = fake_bridge.calls[0] + assert call["api_key"] == "os.environ/MISTRAL_PROVIDER_KEY" + assert call["api_base"] == "os.environ/MISTRAL_PROVIDER_BASE" + + +@pytest.mark.asyncio +async def test_aocr_forwards_os_environ_api_key_reference_to_rust( + fake_async_bridge: RecordingAsyncBridge, +) -> None: + await litellm.aocr( + model=MODEL, document=DOCUMENT, api_key="os.environ/MISTRAL_OCR_TEST_KEY" + ) + + assert fake_async_bridge.calls[0]["api_key"] == "os.environ/MISTRAL_OCR_TEST_KEY"