mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
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:
parent
2cf209f265
commit
4da89566eb
7 changed files with 315 additions and 7 deletions
85
litellm-rust/crates/ai-gateway/src/config.rs
Normal file
85
litellm-rust/crates/ai-gateway/src/config.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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/";
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue