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>
This commit is contained in:
devin-ai-integration[bot] 2026-07-16 21:05:52 -07:00 committed by GitHub
parent 2cf209f265
commit 4da89566eb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 315 additions and 7 deletions

View file

@ -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<String> + Sync),
) -> Option<String> {
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<String> {
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);
}
}

View file

@ -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/";

View file

@ -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<Value> {
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<String> + Sync),
) -> CoreResult<Value> {
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<Value> {
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());

View file

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

View file

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

View file

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

View file

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