mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
wip
This commit is contained in:
parent
592fd00504
commit
6aeae9354a
39 changed files with 919 additions and 902 deletions
|
|
@ -17,7 +17,6 @@ pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Resu
|
|||
.await
|
||||
}
|
||||
|
||||
|
||||
pub fn admit(
|
||||
model: &str,
|
||||
provider: Option<&str>,
|
||||
|
|
|
|||
|
|
@ -55,7 +55,6 @@ pub fn chat_completions_decline_reason(
|
|||
.map(|reason| reason.0)
|
||||
}
|
||||
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, serde::Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct AdmissionContext {
|
||||
|
|
|
|||
|
|
@ -27,7 +27,6 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Re
|
|||
execute_messages_provider_stream(request).await
|
||||
}
|
||||
|
||||
|
||||
pub fn admit(
|
||||
model: &str,
|
||||
provider: Option<&str>,
|
||||
|
|
|
|||
|
|
@ -112,7 +112,11 @@ mod tests {
|
|||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
for (error, status, message) in [
|
||||
(Error::Connect("connection refused".into()), 0, "connection refused"),
|
||||
(
|
||||
Error::Connect("connection refused".into()),
|
||||
0,
|
||||
"connection refused",
|
||||
),
|
||||
(
|
||||
Error::Http {
|
||||
status: 429,
|
||||
|
|
|
|||
|
|
@ -269,7 +269,7 @@ mod tests {
|
|||
(
|
||||
"messages",
|
||||
"amessages",
|
||||
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, has_agentic_hook=None)",
|
||||
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, has_agentic_hook=None, on_request=None)",
|
||||
),
|
||||
(
|
||||
"chat_completions",
|
||||
|
|
|
|||
|
|
@ -26,6 +26,14 @@ fn prepare_messages(
|
|||
options.custom_llm_provider.as_deref(),
|
||||
inputs.has_agentic_hook.unwrap_or(false),
|
||||
))?;
|
||||
if let Some(on_request) = inputs.on_request {
|
||||
Python::attach(|py| {
|
||||
on_request
|
||||
.call0(py)
|
||||
.map(|_| ())
|
||||
.map_err(|error| crate::errors::host_callback_error(py, error))
|
||||
})?;
|
||||
}
|
||||
|
||||
Ok(async move {
|
||||
let RouteOptions {
|
||||
|
|
@ -66,6 +74,7 @@ bridge_route! {
|
|||
extra_headers: Option<serde_json::Value>,
|
||||
timeout_seconds: Option<f64>,
|
||||
has_agentic_hook: Option<bool>,
|
||||
on_request: Option<Py<PyAny>>,
|
||||
},
|
||||
prepare = prepare_messages,
|
||||
errors = execution_error_to_pyerr,
|
||||
|
|
|
|||
|
|
@ -34,8 +34,9 @@ fn count_input_tokens<'py>(
|
|||
legacy_accounting: bool,
|
||||
resource_loader: Py<PyAny>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let tokenizer = litellm_token_counter::admit_tokenizer(kind, encoding, disabled, legacy_accounting)
|
||||
.map_err(admission_error_to_pyerr)?;
|
||||
let tokenizer =
|
||||
litellm_token_counter::admit_tokenizer(kind, encoding, disabled, legacy_accounting)
|
||||
.map_err(admission_error_to_pyerr)?;
|
||||
CoreTokenCounter::admit_request(body).map_err(admission_error_to_pyerr)?;
|
||||
let cached = cached_counter(py, tokenizer, resource_loader)?;
|
||||
let body = body.to_vec();
|
||||
|
|
@ -72,7 +73,8 @@ fn cached_counter(
|
|||
.call1(py, (tokenizer,))
|
||||
.and_then(|value| value.extract(py))
|
||||
.map_err(|error| RustBridgeUnavailable::new_err(error.to_string()))?;
|
||||
let counter = release_gil(py, move || load_counter(tokenizer, &resource)).map_err(token_count_error_to_pyerr)?;
|
||||
let counter = release_gil(py, move || load_counter(tokenizer, &resource))
|
||||
.map_err(token_count_error_to_pyerr)?;
|
||||
let cached = Arc::new(CachedCounter {
|
||||
counter: Arc::new(counter),
|
||||
encode_slots: Arc::new(Semaphore::new(encode_parallelism())),
|
||||
|
|
@ -109,7 +111,9 @@ fn admission_error_to_pyerr(error: Error) -> PyErr {
|
|||
fn token_count_error_to_pyerr(error: Error) -> PyErr {
|
||||
let message = error.to_string();
|
||||
match error {
|
||||
Error::Load(_) | Error::Ranks(_) | Error::UnicodeClasses => RustBridgeUnavailable::new_err(message),
|
||||
Error::Load(_) | Error::Ranks(_) | Error::UnicodeClasses => {
|
||||
RustBridgeUnavailable::new_err(message)
|
||||
}
|
||||
Error::UnsupportedTokenizer
|
||||
| Error::RequestParse(_)
|
||||
| Error::MissingInput
|
||||
|
|
|
|||
|
|
@ -449,6 +449,83 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
timeout=timeout,
|
||||
)
|
||||
|
||||
def completion_dispatch() -> "ModelResponse | CustomStreamWrapper":
|
||||
request_headers, data = finish_request(
|
||||
config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=transform_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
)
|
||||
if stream is True:
|
||||
data["stream"] = stream
|
||||
completion_stream, response_headers = make_sync_call(
|
||||
client=client,
|
||||
api_base=api_base,
|
||||
headers=request_headers,
|
||||
data=json.dumps(data),
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
timeout=timeout,
|
||||
json_mode=json_mode,
|
||||
speed=optional_params.get("speed") if optional_params else None,
|
||||
tool_name_reverse_map=(
|
||||
litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY)
|
||||
if isinstance(litellm_params, dict)
|
||||
else None
|
||||
),
|
||||
)
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="anthropic",
|
||||
logging_obj=logging_obj,
|
||||
_response_headers=process_anthropic_headers(response_headers),
|
||||
)
|
||||
|
||||
sync_client: Final = (
|
||||
client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout})
|
||||
)
|
||||
try:
|
||||
response: Final = sync_client.post(
|
||||
api_base,
|
||||
headers=request_headers,
|
||||
data=json.dumps(data),
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
status_code: Final = getattr(e, "status_code", 500)
|
||||
error_headers = getattr(e, "headers", None)
|
||||
error_text = getattr(e, "text", str(e))
|
||||
error_response: Final[object] = getattr(e, "response", None)
|
||||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
if error_response and hasattr(error_response, "text"):
|
||||
error_text = getattr(error_response, "text", error_text)
|
||||
raise AnthropicError(
|
||||
message=error_text,
|
||||
status_code=status_code,
|
||||
headers=error_headers,
|
||||
)
|
||||
|
||||
return config.transform_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
model_response=model_response,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
request_data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
encoding=encoding,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
|
||||
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
|
||||
**AnthropicConfig.get_config(model=model),
|
||||
**optional_params,
|
||||
|
|
@ -462,9 +539,10 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
}
|
||||
log_rust_pre_call: Final = lambda: logging_obj.pre_call(
|
||||
input=messages, api_key=api_key, additional_args=rust_logging_args
|
||||
)
|
||||
|
||||
def log_rust_pre_call() -> None:
|
||||
logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args)
|
||||
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
|
|
@ -488,7 +566,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
on_response=log_rust_post_call,
|
||||
python_fallback=acompletion_dispatch,
|
||||
)
|
||||
rust_response: Final = rust_chat_completions_bridge.chat_completions(
|
||||
return rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
|
|
@ -502,94 +580,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
litellm_params=litellm_params,
|
||||
on_request=log_rust_pre_call,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=lambda: None,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
|
||||
if acompletion is True:
|
||||
return acompletion_dispatch()
|
||||
else:
|
||||
headers, data = finish_request(
|
||||
config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=transform_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
)
|
||||
## COMPLETION CALL
|
||||
if (
|
||||
stream is True
|
||||
): # if function call - fake the streaming (need complete blocks for output parsing in openai format)
|
||||
data["stream"] = stream
|
||||
completion_stream, headers = make_sync_call(
|
||||
client=client,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
timeout=timeout,
|
||||
json_mode=json_mode,
|
||||
speed=optional_params.get("speed") if optional_params else None,
|
||||
tool_name_reverse_map=(
|
||||
litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY)
|
||||
if isinstance(litellm_params, dict)
|
||||
else None
|
||||
),
|
||||
)
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="anthropic",
|
||||
logging_obj=logging_obj,
|
||||
_response_headers=process_anthropic_headers(headers),
|
||||
)
|
||||
|
||||
else:
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
client = _get_httpx_client(params={"timeout": timeout})
|
||||
else:
|
||||
client = client
|
||||
|
||||
try:
|
||||
response: Final = client.post(
|
||||
api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
status_code: Final = getattr(e, "status_code", 500)
|
||||
error_headers = getattr(e, "headers", None)
|
||||
error_text = getattr(e, "text", str(e))
|
||||
error_response: Final[object] = getattr(e, "response", None)
|
||||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
if error_response and hasattr(error_response, "text"):
|
||||
error_text = getattr(error_response, "text", error_text)
|
||||
raise AnthropicError(
|
||||
message=error_text,
|
||||
status_code=status_code,
|
||||
headers=error_headers,
|
||||
)
|
||||
|
||||
return config.transform_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
model_response=model_response,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
request_data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
encoding=encoding,
|
||||
json_mode=json_mode,
|
||||
python_fallback=completion_dispatch,
|
||||
)
|
||||
|
||||
def embedding(self):
|
||||
|
|
|
|||
|
|
@ -50,9 +50,8 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
python_fallback=None,
|
||||
)
|
||||
if rust_response is None:
|
||||
raise RuntimeError("Rust audio transcription bridge is unavailable")
|
||||
return TranscriptionResponse(**rust_response)
|
||||
|
||||
async def async_audio_transcriptions(
|
||||
|
|
@ -76,7 +75,6 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
python_fallback=None,
|
||||
)
|
||||
if rust_response is None:
|
||||
raise RuntimeError("Rust audio transcription bridge is unavailable")
|
||||
return TranscriptionResponse(**rust_response)
|
||||
|
|
|
|||
|
|
@ -411,49 +411,139 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
"api_base": proxy_endpoint_url,
|
||||
"headers": headers,
|
||||
}
|
||||
log_rust_pre_call: Final = lambda: logging_obj.pre_call(
|
||||
input=messages, api_key="", additional_args=rust_logging_args
|
||||
)
|
||||
|
||||
def log_rust_pre_call() -> None:
|
||||
logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args)
|
||||
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
api_key="",
|
||||
additional_args=rust_logging_args,
|
||||
)
|
||||
if acompletion:
|
||||
return rust_chat_completions_bridge.achat_completions(
|
||||
|
||||
def completion_dispatch() -> ModelResponse | CustomStreamWrapper:
|
||||
request_data: Final = litellm.AmazonConverseConfig()._transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=extra_headers,
|
||||
)
|
||||
data: Final = json.dumps(request_data)
|
||||
prepped: Final = self.get_request_headers(
|
||||
credentials=credentials,
|
||||
aws_region_name=aws_region_name,
|
||||
extra_headers=extra_headers,
|
||||
endpoint_url=proxy_endpoint_url,
|
||||
data=data,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
client_timeout: Final = httpx.Timeout(timeout) if isinstance(timeout, (float, int)) else timeout
|
||||
sync_client: Final = (
|
||||
client
|
||||
if isinstance(client, HTTPHandler)
|
||||
else _get_httpx_client({} if client_timeout is None else {"timeout": client_timeout})
|
||||
)
|
||||
if stream is True:
|
||||
completion_stream, response_headers = make_sync_call(
|
||||
client=sync_client,
|
||||
api_base=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
logging_obj=logging_obj,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
_response_headers=response_headers,
|
||||
)
|
||||
|
||||
try:
|
||||
response: Final = sync_client.post(
|
||||
url=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(
|
||||
status_code=error_code,
|
||||
message=error_response_text(err.response),
|
||||
headers=err.response.headers,
|
||||
response=err.response,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
|
||||
model=model,
|
||||
response=response,
|
||||
model_response=model_response,
|
||||
stream=stream if isinstance(stream, bool) else False,
|
||||
logging_obj=logging_obj,
|
||||
api_key="",
|
||||
data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
encoding=encoding,
|
||||
)
|
||||
transformed_response.set_provider_response_headers(response.headers)
|
||||
return transformed_response
|
||||
|
||||
if acompletion:
|
||||
return rust_chat_completions_bridge.achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
on_request=log_rust_pre_call,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=lambda: self.async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=proxy_endpoint_url,
|
||||
model_response=model_response,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
on_request=log_rust_pre_call,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=lambda: self.async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=proxy_endpoint_url,
|
||||
model_response=model_response,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
credentials=credentials,
|
||||
api_key=api_key,
|
||||
),
|
||||
)
|
||||
rust_response: Final = rust_chat_completions_bridge.chat_completions(
|
||||
logger_fn=logger_fn,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
credentials=credentials,
|
||||
api_key=api_key,
|
||||
),
|
||||
)
|
||||
return rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
|
|
@ -467,151 +557,5 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
litellm_params=litellm_params,
|
||||
on_request=log_rust_pre_call,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=lambda: None,
|
||||
python_fallback=completion_dispatch,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
|
||||
### ROUTING (ASYNC, STREAMING, SYNC)
|
||||
if acompletion:
|
||||
if isinstance(client, HTTPHandler):
|
||||
client = None
|
||||
if stream is True:
|
||||
return self.async_streaming(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=proxy_endpoint_url,
|
||||
model_response=model_response,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=True,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
credentials=credentials,
|
||||
api_key=api_key,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
### ASYNC COMPLETION
|
||||
return self.async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=proxy_endpoint_url,
|
||||
model_response=model_response,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
credentials=credentials,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
## TRANSFORMATION ##
|
||||
|
||||
_data: Final = litellm.AmazonConverseConfig()._transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=extra_headers,
|
||||
)
|
||||
data: Final = json.dumps(_data)
|
||||
|
||||
prepped: Final = self.get_request_headers(
|
||||
credentials=credentials,
|
||||
aws_region_name=aws_region_name,
|
||||
extra_headers=extra_headers,
|
||||
endpoint_url=proxy_endpoint_url,
|
||||
data=data,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
if client is None or isinstance(client, AsyncHTTPHandler):
|
||||
_params: Final = {}
|
||||
if timeout is not None:
|
||||
if isinstance(timeout, float) or isinstance(timeout, int):
|
||||
timeout = httpx.Timeout(timeout)
|
||||
_params["timeout"] = timeout
|
||||
client = _get_httpx_client(_params)
|
||||
else:
|
||||
client = client
|
||||
|
||||
if stream is not None and stream is True:
|
||||
completion_stream, response_headers = make_sync_call(
|
||||
client=(client if client is not None and isinstance(client, HTTPHandler) else None),
|
||||
api_base=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
streaming_response: Final = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
_response_headers=response_headers,
|
||||
)
|
||||
|
||||
return streaming_response
|
||||
|
||||
### COMPLETION
|
||||
|
||||
try:
|
||||
response: Final = client.post(
|
||||
url=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(
|
||||
status_code=error_code,
|
||||
message=error_response_text(err.response),
|
||||
headers=err.response.headers,
|
||||
response=err.response,
|
||||
)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
|
||||
model=model,
|
||||
response=response,
|
||||
model_response=model_response,
|
||||
stream=stream if isinstance(stream, bool) else False,
|
||||
logging_obj=logging_obj,
|
||||
api_key="",
|
||||
data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
encoding=encoding,
|
||||
)
|
||||
sync_transformed_response.set_provider_response_headers(response.headers)
|
||||
return sync_transformed_response
|
||||
|
|
|
|||
|
|
@ -2264,36 +2264,90 @@ class BaseLLMHTTPHandler:
|
|||
# internally -- this only deduplicates the success path.
|
||||
request_body_json: Final = json.dumps(request_body)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=[{"role": "user", "content": request_body_json}],
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": request_body,
|
||||
"api_base": str(request_url),
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
def log_pre_call() -> None:
|
||||
logging_obj.pre_call(
|
||||
input=[{"role": "user", "content": request_body_json}],
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": request_body,
|
||||
"api_base": str(request_url),
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
rust_messages_response: Final = await self._maybe_rust_anthropic_messages(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
|
||||
model=model,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
request_body=request_body,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
async def python_fallback() -> AnthropicMessagesResponse | AsyncIterator:
|
||||
log_pre_call()
|
||||
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
|
||||
async_httpx_client=async_httpx_client,
|
||||
request_url=request_url,
|
||||
headers=headers,
|
||||
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
|
||||
request_body=request_body,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
if rust_messages_response is not None:
|
||||
logging_obj=logging_obj,
|
||||
provider_config=anthropic_messages_provider_config,
|
||||
litellm_params=litellm_params,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
logging_obj.model_call_details["httpx_response"] = response
|
||||
if stream:
|
||||
return self._rust_anthropic_messages_fake_stream(rust_messages_response)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
AnthropicMessagesStreamingResponse,
|
||||
anthropic_messages_stream_hidden_params,
|
||||
)
|
||||
|
||||
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
|
||||
model=model,
|
||||
httpx_response=response,
|
||||
request_body=request_body,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
|
||||
if not self._has_agentic_completion_hook(logging_obj):
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=completion_stream,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
AgenticAnthropicStreamingIterator,
|
||||
)
|
||||
|
||||
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
|
||||
logging_obj=logging_obj,
|
||||
tools=anthropic_messages_optional_request_params.get("tools"),
|
||||
)
|
||||
agentic_stream: Final = AgenticAnthropicStreamingIterator(
|
||||
completion_stream=completion_stream,
|
||||
http_handler=self,
|
||||
model=model,
|
||||
messages=messages,
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
|
||||
hold_back=bool(held_back_tool_names),
|
||||
server_fulfilled_tool_names=held_back_tool_names,
|
||||
)
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=agentic_stream,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
|
||||
initial_response: Final = anthropic_messages_provider_config.transform_anthropic_messages_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return await self._finalize_anthropic_messages_response(
|
||||
initial_response=rust_messages_response,
|
||||
initial_response=initial_response,
|
||||
model=model,
|
||||
messages=messages,
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
|
|
@ -2304,96 +2358,42 @@ class BaseLLMHTTPHandler:
|
|||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
|
||||
async_httpx_client=async_httpx_client,
|
||||
request_url=request_url,
|
||||
headers=headers,
|
||||
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
|
||||
request_body=request_body,
|
||||
stream=stream or False,
|
||||
logging_obj=logging_obj,
|
||||
provider_config=anthropic_messages_provider_config,
|
||||
litellm_params=litellm_params,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
|
||||
# used for logging + cost tracking
|
||||
logging_obj.model_call_details["httpx_response"] = response
|
||||
|
||||
initial_response: AsyncIterator | AnthropicMessagesResponse
|
||||
if stream:
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
AnthropicMessagesStreamingResponse,
|
||||
anthropic_messages_stream_hidden_params,
|
||||
)
|
||||
|
||||
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
|
||||
model=model,
|
||||
httpx_response=response,
|
||||
request_body=request_body,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
|
||||
|
||||
if not self._has_agentic_completion_hook(logging_obj):
|
||||
# No callback overrides async_should_run_agentic_loop, so the
|
||||
# agentic wrapper's only effect would be buffering every chunk
|
||||
# and rebuilding the response from SSE at end-of-stream to call
|
||||
# hooks that all return (False, {}). Stream through directly and
|
||||
# skip that per-chunk + end-of-stream overhead.
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=completion_stream,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
AgenticAnthropicStreamingIterator,
|
||||
)
|
||||
|
||||
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
|
||||
logging_obj=logging_obj,
|
||||
tools=anthropic_messages_optional_request_params.get("tools"),
|
||||
)
|
||||
initial_response = AgenticAnthropicStreamingIterator(
|
||||
completion_stream=completion_stream,
|
||||
http_handler=self,
|
||||
async def adapt_rust_response(response: dict[str, object]) -> AnthropicMessagesResponse | AsyncIterator:
|
||||
response_obj: Final = cast(AnthropicMessagesResponse, dict(response))
|
||||
response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
|
||||
if stream:
|
||||
return self._rust_anthropic_messages_fake_stream(response_obj)
|
||||
return await self._finalize_anthropic_messages_response(
|
||||
initial_response=response_obj,
|
||||
model=model,
|
||||
messages=messages,
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
|
||||
hold_back=bool(held_back_tool_names),
|
||||
server_fulfilled_tool_names=held_back_tool_names,
|
||||
)
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=initial_response,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
else:
|
||||
initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
return await self._finalize_anthropic_messages_response(
|
||||
initial_response=initial_response,
|
||||
from litellm.rust_bridge import messages as rust_messages_bridge
|
||||
|
||||
upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"}
|
||||
return await rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
messages=messages,
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
body=upstream_body,
|
||||
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
|
||||
api_key=api_key,
|
||||
kwargs=kwargs,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
on_request=log_pre_call,
|
||||
python_fallback=python_fallback,
|
||||
adapt=adapt_rust_response,
|
||||
)
|
||||
|
||||
async def _finalize_anthropic_messages_response(
|
||||
|
|
@ -2432,39 +2432,6 @@ class BaseLLMHTTPHandler:
|
|||
"anthropic_messages",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _maybe_rust_anthropic_messages(
|
||||
*,
|
||||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
has_agentic_hook: bool,
|
||||
model: str,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
headers: dict,
|
||||
request_body: dict,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> AnthropicMessagesResponse | None:
|
||||
from litellm.rust_bridge import messages as rust_messages_bridge
|
||||
|
||||
upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"}
|
||||
rust_response: Final = await rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
body=upstream_body,
|
||||
has_agentic_hook=has_agentic_hook,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_response is None:
|
||||
return None
|
||||
|
||||
response_obj: Final = cast(AnthropicMessagesResponse, dict(rust_response))
|
||||
response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
|
||||
return response_obj
|
||||
|
||||
@staticmethod
|
||||
def _rust_anthropic_messages_fake_stream(
|
||||
rust_response: AnthropicMessagesResponse,
|
||||
|
|
@ -6635,27 +6602,26 @@ class BaseLLMHTTPHandler:
|
|||
async def _backend_connection():
|
||||
from litellm.rust_bridge.responses import websocket as rust_responses_websocket
|
||||
|
||||
rust_backend: Final = await rust_responses_websocket.connect(
|
||||
async def python_fallback() -> ClientConnection:
|
||||
return await websockets.connect(
|
||||
ws_url,
|
||||
additional_headers=headers,
|
||||
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
||||
ssl=ssl_context,
|
||||
)
|
||||
|
||||
backend: Final = await rust_responses_websocket.connect(
|
||||
url=ws_url,
|
||||
headers={str(key): str(value) for key, value in headers.items()},
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
python_fallback=python_fallback,
|
||||
)
|
||||
if rust_backend is not None:
|
||||
try:
|
||||
yield rust_backend
|
||||
finally:
|
||||
await rust_backend.close()
|
||||
return
|
||||
|
||||
async with websockets.connect(
|
||||
ws_url,
|
||||
additional_headers=headers,
|
||||
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
||||
ssl=ssl_context,
|
||||
) as backend:
|
||||
try:
|
||||
yield backend
|
||||
finally:
|
||||
await backend.close()
|
||||
|
||||
async with _backend_connection() as backend_ws:
|
||||
_request_data: Final[dict[str, object]] = {}
|
||||
|
|
|
|||
|
|
@ -80,14 +80,18 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr
|
|||
native_call: Final[Callable[[], Awaitable[OCRResponse]] | None] = (
|
||||
(lambda: native(request, args, kwargs, True, HOST)) if native is not None else None
|
||||
)
|
||||
|
||||
async def python_fallback() -> OCRResponse:
|
||||
return await fallback(*args, **kwargs)
|
||||
|
||||
async def adapt(value: OCRResponse) -> OCRResponse:
|
||||
return value
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda value: value,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(
|
||||
route=COMPONENT.name.value,
|
||||
provider=request.custom_llm_provider or "",
|
||||
|
|
|
|||
|
|
@ -1392,50 +1392,34 @@ async def count_request_input_tokens(
|
|||
|
||||
Tokenizing is the reservation path's dominant CPU cost and is O(prompt), so
|
||||
counting a large prompt inline stalls every other request on the worker.
|
||||
Models whose tokenizer the Rust bridge ports (Anthropic, tiktoken cl100k_base
|
||||
and o200k_base) are counted from the raw body by the bridge when it is enabled, once per
|
||||
distinct tokenizer, which parses and tokenizes with the GIL released.
|
||||
Everything it declines is counted in Python, large prompts in a worker
|
||||
thread. The counts are reused by both the max-cost and the input-cost
|
||||
estimate.
|
||||
The Rust bridge counts the complete request once per distinct tokenizer.
|
||||
If that operation is unavailable or declines any tokenizer, the complete
|
||||
request is counted in Python, with large prompts moved to a worker thread.
|
||||
The counts are reused by both the max-cost and input-cost estimates.
|
||||
"""
|
||||
models: Final = _get_request_models(request_body=request_body, route=route, llm_router=llm_router)
|
||||
if not models:
|
||||
return MappingProxyType({})
|
||||
tokenizers: Final[Mapping[str, RustTokenizer | None]] = MappingProxyType(
|
||||
tokenizers: Final[Mapping[str, RustTokenizer]] = MappingProxyType(
|
||||
{model: rust_tokenizer(model) for model in models}
|
||||
)
|
||||
distinct_tokenizers: Final[tuple[RustTokenizer, ...]] = tuple(
|
||||
dict.fromkeys(tokenizer for tokenizer in tokenizers.values() if tokenizer is not None)
|
||||
)
|
||||
rust_counts_by_tokenizer: Final[Mapping[RustTokenizer, int]] = MappingProxyType(
|
||||
{
|
||||
tokenizer: count.input_tokens
|
||||
for tokenizer in distinct_tokenizers
|
||||
if raw_body is not None and (count := await count_input_tokens(raw_body, tokenizer)) is not None
|
||||
}
|
||||
)
|
||||
rust_counts: Final = MappingProxyType(
|
||||
{
|
||||
model: rust_counts_by_tokenizer[tokenizer]
|
||||
for model, tokenizer in tokenizers.items()
|
||||
if tokenizer is not None and tokenizer in rust_counts_by_tokenizer
|
||||
}
|
||||
)
|
||||
python_models: Final = tuple(model for model in models if model not in rust_counts)
|
||||
python_counts: Final = (
|
||||
MappingProxyType({})
|
||||
if not python_models
|
||||
else _count_input_tokens_for_models(request_body=request_body, models=python_models)
|
||||
if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS
|
||||
else await asyncio.to_thread(
|
||||
|
||||
async def python_fallback() -> Mapping[str, int]:
|
||||
if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS:
|
||||
return _count_input_tokens_for_models(request_body=request_body, models=models)
|
||||
return await asyncio.to_thread(
|
||||
_count_input_tokens_for_models,
|
||||
request_body=request_body,
|
||||
models=python_models,
|
||||
models=models,
|
||||
)
|
||||
|
||||
counts: Final = await count_input_tokens(
|
||||
body=raw_body,
|
||||
tokenizers=tokenizers,
|
||||
python_fallback=python_fallback,
|
||||
)
|
||||
verbose_proxy_logger.debug("input token counts: rust=%s python=%s", dict(rust_counts), dict(python_counts))
|
||||
return MappingProxyType({**rust_counts, **python_counts})
|
||||
verbose_proxy_logger.debug("input token counts: %s", dict(counts))
|
||||
return counts
|
||||
|
||||
|
||||
def _count_input_tokens_for_models(
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ Optional capabilities use the `litellm.rust(bool)` process override first, `LITE
|
|||
|
||||
## Fallback contract
|
||||
|
||||
Each API calls its native entrypoint at most once. Rust performs request admission inside that entrypoint before provider calls or host callbacks. An unavailable binding or `RustBridgeDeclined` selects the supplied Python fallback only when the policy allows it
|
||||
Each API calls its native entrypoint at most once. Rust performs request admission inside that entrypoint before provider calls or host callbacks. An unavailable binding or `RustBridgeDeclined` permits fallback only when the resolved capability declares Python availability and the caller supplies a real Python callback. Python-capable decisions without a callback, and Rust-required decisions with one, are contract errors
|
||||
|
||||
Provider failures, host callback failures, cancellation, conversion failures, and response adaptation failures propagate without replay. Adaptation runs outside the decline-catching boundary
|
||||
|
||||
|
|
|
|||
|
|
@ -109,6 +109,7 @@ def messages(
|
|||
extra_headers: object = None,
|
||||
timeout_seconds: float | None = None,
|
||||
has_agentic_hook: bool | None = None,
|
||||
on_request: Callable[[], None] | None = None,
|
||||
) -> dict[str, object]: ...
|
||||
def amessages(
|
||||
model: str,
|
||||
|
|
@ -119,6 +120,7 @@ def amessages(
|
|||
extra_headers: object = None,
|
||||
timeout_seconds: float | None = None,
|
||||
has_agentic_hook: bool | None = None,
|
||||
on_request: Callable[[], None] | None = None,
|
||||
) -> Future[dict[str, object]]: ...
|
||||
def chat_completions(
|
||||
model: str,
|
||||
|
|
@ -175,38 +177,37 @@ def count_input_tokens(
|
|||
legacy_accounting: bool,
|
||||
resource_loader: Callable[[str], str],
|
||||
) -> Future[dict[str, object]]: ...
|
||||
|
||||
def gil_stats() -> dict[str, int]: ...
|
||||
|
||||
__all__ = [
|
||||
"RustBridgeUnavailable",
|
||||
"_OCR_MAX_FILE_BYTES",
|
||||
"ResponsesWebSocketConnection",
|
||||
"RustBridgeDeclined",
|
||||
"RustBridgeUnavailable",
|
||||
"RustHostCallbackError",
|
||||
"RustUpstreamError",
|
||||
"ocr",
|
||||
"aocr",
|
||||
"_OCR_MAX_FILE_BYTES",
|
||||
"_ocr_upload_document",
|
||||
"_ocr_file_document",
|
||||
"_ocr_mime_type",
|
||||
"_ocr_lifecycle",
|
||||
"_transcription_lifecycle",
|
||||
"transcription",
|
||||
"atranscription",
|
||||
"_messages_lifecycle",
|
||||
"messages",
|
||||
"amessages",
|
||||
"_chat_completions_lifecycle",
|
||||
"chat_completions",
|
||||
"achat_completions",
|
||||
"_embeddings_lifecycle",
|
||||
"_image_edit_lifecycle",
|
||||
"_image_generation_lifecycle",
|
||||
"_messages_lifecycle",
|
||||
"_moderation_lifecycle",
|
||||
"_ocr_file_document",
|
||||
"_ocr_lifecycle",
|
||||
"_ocr_mime_type",
|
||||
"_ocr_upload_document",
|
||||
"_rerank_lifecycle",
|
||||
"ResponsesWebSocketConnection",
|
||||
"_responses_lifecycle",
|
||||
"_speech_lifecycle",
|
||||
"_transcription_lifecycle",
|
||||
"achat_completions",
|
||||
"amessages",
|
||||
"aocr",
|
||||
"atranscription",
|
||||
"chat_completions",
|
||||
"count_input_tokens",
|
||||
"gil_stats",
|
||||
"messages",
|
||||
"ocr",
|
||||
"transcription",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -90,7 +91,7 @@ def _component(
|
|||
return NativeComponent(name=name, capability=capability, exports=exports)
|
||||
|
||||
|
||||
COMPONENTS: Final = MappingProxyType(
|
||||
COMPONENTS: Final[Mapping[ComponentName, NativeComponent]] = MappingProxyType(
|
||||
{
|
||||
ComponentName.OCR: _component(
|
||||
ComponentName.OCR,
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ import httpx
|
|||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
convert_to_model_response_object,
|
||||
convert_to_model_response_object, # pyright: ignore[reportUnknownVariableType] # legacy converter lacks complete annotations
|
||||
)
|
||||
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
|
||||
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
|
||||
|
|
@ -112,8 +112,9 @@ def chat_completions(
|
|||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
native_call: Final[Callable[[], Mapping[str, object]] | None] = (
|
||||
lambda: rust_chat_completions(
|
||||
def native_call() -> Mapping[str, object]:
|
||||
assert rust_chat_completions is not None
|
||||
return rust_chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -125,12 +126,10 @@ def chat_completions(
|
|||
host_facts=_host_facts(stream, litellm_params),
|
||||
on_request=on_request,
|
||||
)
|
||||
if rust_chat_completions is not None
|
||||
else None
|
||||
)
|
||||
|
||||
return invoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
native_call=native_call if rust_chat_completions is not None else None,
|
||||
python_fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
|
|
@ -163,12 +162,13 @@ async def achat_completions(
|
|||
)
|
||||
rust_achat_completions: Final = execution.select(_ACHAT)
|
||||
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
async def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
native_call: Final[Callable[[], Awaitable[Mapping[str, object]]] | None] = (
|
||||
lambda: rust_achat_completions(
|
||||
def native_call() -> Awaitable[Mapping[str, object]]:
|
||||
assert rust_achat_completions is not None
|
||||
return rust_achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -180,12 +180,10 @@ async def achat_completions(
|
|||
host_facts=_host_facts(stream, litellm_params),
|
||||
on_request=on_request,
|
||||
)
|
||||
if rust_achat_completions is not None
|
||||
else None
|
||||
)
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
native_call=native_call if rust_achat_completions is not None else None,
|
||||
python_fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Mapping
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import Protocol
|
||||
|
||||
|
||||
|
|
@ -15,6 +15,7 @@ class RustMessages(Protocol):
|
|||
extra_headers: Mapping[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
has_agentic_hook: bool = False,
|
||||
on_request: Callable[[], None] | None = None,
|
||||
) -> dict[str, object]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
|
@ -30,5 +31,6 @@ class RustAmessages(Protocol):
|
|||
extra_headers: Mapping[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
has_agentic_hook: bool = False,
|
||||
on_request: Callable[[], None] | None = None,
|
||||
) -> Awaitable[dict[str, object]]:
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from __future__ import annotations
|
|||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import (
|
||||
Final,
|
||||
TypeVar,
|
||||
cast, # noqa: TID251 # native callable signatures are checked by bridge contract tests
|
||||
)
|
||||
|
||||
|
|
@ -28,6 +29,7 @@ def _as_amessages(value: object) -> RustAmessages | None:
|
|||
|
||||
_MESSAGES: Final = COMPONENT.bind("messages", validate=_as_messages)
|
||||
_AMESSAGES: Final = COMPONENT.bind("amessages", validate=_as_amessages)
|
||||
ResultT = TypeVar("ResultT")
|
||||
|
||||
|
||||
def set_rust_messages(
|
||||
|
|
@ -57,7 +59,10 @@ def messages(
|
|||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> dict[str, object] | None:
|
||||
on_request: Callable[[], None] = lambda: None,
|
||||
python_fallback: Callable[[], ResultT],
|
||||
adapt: Callable[[dict[str, object]], ResultT],
|
||||
) -> ResultT:
|
||||
execution: Final = COMPONENT.resolve(
|
||||
CapabilityContext(
|
||||
provider=custom_llm_provider or "",
|
||||
|
|
@ -77,6 +82,7 @@ def messages(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
on_request=on_request,
|
||||
)
|
||||
)
|
||||
if rust_messages is not None
|
||||
|
|
@ -85,8 +91,8 @@ def messages(
|
|||
return invoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=lambda: None,
|
||||
adapt=lambda value: value,
|
||||
python_fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(
|
||||
route=COMPONENT.name.value,
|
||||
provider=custom_llm_provider or "",
|
||||
|
|
@ -105,7 +111,10 @@ async def amessages(
|
|||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> dict[str, object] | None:
|
||||
on_request: Callable[[], None] = lambda: None,
|
||||
python_fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[dict[str, object]], Awaitable[ResultT]],
|
||||
) -> ResultT:
|
||||
execution: Final = COMPONENT.resolve(
|
||||
CapabilityContext(
|
||||
provider=custom_llm_provider or "",
|
||||
|
|
@ -125,20 +134,18 @@ async def amessages(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
on_request=on_request,
|
||||
)
|
||||
)
|
||||
if rust_amessages is not None
|
||||
else None
|
||||
)
|
||||
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda value: value,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(
|
||||
route=COMPONENT.name.value,
|
||||
provider=custom_llm_provider or "",
|
||||
|
|
|
|||
|
|
@ -1,21 +1,9 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.ocr.definition import COMPONENT
|
||||
from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest, RustAocr, RustOcr
|
||||
from litellm.rust_bridge.ocr.value import (
|
||||
aocr,
|
||||
load_rust_aocr,
|
||||
load_rust_ocr,
|
||||
ocr,
|
||||
)
|
||||
from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest
|
||||
|
||||
__all__: Final = (
|
||||
"COMPONENT",
|
||||
"LiteLLMOcrRequest",
|
||||
"RustAocr",
|
||||
"RustOcr",
|
||||
"aocr",
|
||||
"load_rust_aocr",
|
||||
"load_rust_ocr",
|
||||
"ocr",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, Protocol, cast
|
||||
from typing import Final, Protocol
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
|
|
@ -35,7 +35,7 @@ class OcrLifecycleHost:
|
|||
request: LiteLLMOcrRequest,
|
||||
request_provider: str,
|
||||
) -> Exception:
|
||||
mapper: Final = cast(ExceptionMapper, litellm.exception_type) # cast-ok: legacy public exception mapper
|
||||
mapper: Final[ExceptionMapper] = litellm.exception_type # pyright: ignore[reportAssignmentType] # legacy public mapper is callable
|
||||
try:
|
||||
return mapper(
|
||||
model=request.model.removeprefix(f"{request_provider}/"),
|
||||
|
|
|
|||
|
|
@ -1,8 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Mapping
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -18,35 +17,3 @@ class LiteLLMOcrRequest:
|
|||
extra_headers: dict[str, object] | None
|
||||
kwargs: Mapping[str, object]
|
||||
input_sources: Mapping[str, str] | None = None
|
||||
|
||||
|
||||
class RustOcr(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
input_sources: dict[str, str],
|
||||
timeout_seconds: float | None,
|
||||
) -> dict[str, object]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class RustAocr(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
input_sources: dict[str, str],
|
||||
timeout_seconds: float | None,
|
||||
) -> Awaitable[dict[str, object]]:
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
|
|
@ -4,35 +4,9 @@ from __future__ import annotations
|
|||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final, cast # noqa: TID251 # native extension exposes dynamically typed callables
|
||||
|
||||
import httpx
|
||||
from typing import Final
|
||||
|
||||
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
|
||||
from litellm.rust_bridge.ocr.definition import COMPONENT
|
||||
from litellm.rust_bridge.ocr.types import RustAocr, RustOcr
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
|
||||
|
||||
|
||||
def _as_ocr(value: object) -> RustOcr | None:
|
||||
return cast(RustOcr, value) if callable(value) else None
|
||||
|
||||
|
||||
def _as_aocr(value: object) -> RustAocr | None:
|
||||
return cast(RustAocr, value) if callable(value) else None
|
||||
|
||||
|
||||
_OCR: Final = COMPONENT.bind("ocr", validate=_as_ocr)
|
||||
_AOCR: Final = COMPONENT.bind("aocr", validate=_as_aocr)
|
||||
|
||||
|
||||
def load_rust_ocr() -> RustOcr | None:
|
||||
return COMPONENT.resolve().select(_OCR)
|
||||
|
||||
|
||||
def load_rust_aocr() -> RustAocr | None:
|
||||
return COMPONENT.resolve().select(_AOCR)
|
||||
|
||||
|
||||
def adapt_response(response: Mapping[str, object]) -> OCRResponse:
|
||||
|
|
@ -43,80 +17,3 @@ def adapt_response(response: Mapping[str, object]) -> OCRResponse:
|
|||
if isinstance(provider_native_response, Mapping):
|
||||
normalized.set_provider_native_response(provider_native_response)
|
||||
return normalized
|
||||
|
||||
|
||||
def ocr(
|
||||
*,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
input_sources: Mapping[str, str] | None = None,
|
||||
) -> dict[str, object] | None:
|
||||
execution: Final = COMPONENT.resolve()
|
||||
rust_ocr: Final = execution.select(_OCR)
|
||||
return invoke(
|
||||
execution=execution,
|
||||
native_call=(
|
||||
lambda: rust_ocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
if rust_ocr is not None
|
||||
else None,
|
||||
python_fallback=lambda: None,
|
||||
adapt=lambda value: value,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
async def aocr(
|
||||
*,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
input_sources: Mapping[str, str] | None = None,
|
||||
) -> dict[str, object] | None:
|
||||
execution: Final = COMPONENT.resolve()
|
||||
rust_aocr: Final = execution.select(_AOCR)
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=(
|
||||
lambda: rust_aocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
if rust_aocr is not None
|
||||
else None,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda value: value,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Final, Protocol, cast # noqa: TID251 # native class is validated at load time
|
||||
from typing import Final, Protocol, TypeVar, cast # noqa: TID251 # native class is validated at load time
|
||||
|
||||
import httpx
|
||||
from websockets.exceptions import ConnectionClosedOK
|
||||
|
|
@ -23,6 +23,14 @@ class RustResponsesWebSocket(Protocol):
|
|||
def close(self) -> Awaitable[None]: ...
|
||||
|
||||
|
||||
class ResponsesWebSocket(Protocol):
|
||||
def send(self, text: str) -> Awaitable[None]: ...
|
||||
|
||||
def recv(self) -> Awaitable[str | bytes]: ...
|
||||
|
||||
def close(self) -> Awaitable[None]: ...
|
||||
|
||||
|
||||
class RustResponsesWebSocketConnection(Protocol):
|
||||
@classmethod
|
||||
def connect(
|
||||
|
|
@ -71,6 +79,9 @@ class _ConnectionAdapter:
|
|||
await self._connection.close()
|
||||
|
||||
|
||||
ConnectionT = TypeVar("ConnectionT", bound=ResponsesWebSocket)
|
||||
|
||||
|
||||
async def connect(
|
||||
*,
|
||||
url: str,
|
||||
|
|
@ -78,7 +89,8 @@ async def connect(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
custom_llm_provider: str | None,
|
||||
model: str,
|
||||
) -> _ConnectionAdapter | None:
|
||||
python_fallback: Callable[[], Awaitable[ConnectionT]],
|
||||
) -> ResponsesWebSocket:
|
||||
execution: Final = COMPONENT.resolve(
|
||||
CapabilityContext(
|
||||
provider=custom_llm_provider or "",
|
||||
|
|
@ -100,16 +112,14 @@ async def connect(
|
|||
if connection_type is not None
|
||||
else None
|
||||
)
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
|
||||
connection: Final = await ainvoke(
|
||||
async def adapt(connection: RustResponsesWebSocket) -> ResponsesWebSocket:
|
||||
return _ConnectionAdapter(connection)
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
native_call=native_call,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda value: value,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
if connection is None:
|
||||
return None
|
||||
return _ConnectionAdapter(connection)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ from __future__ import annotations
|
|||
from collections.abc import Callable, Coroutine
|
||||
from dataclasses import dataclass
|
||||
from types import ModuleType
|
||||
from typing import Final, Literal, Protocol, TypeVar, cast, overload
|
||||
from typing import Final, Literal, Protocol, TypeVar, overload
|
||||
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.configuration import (
|
||||
|
|
@ -55,7 +55,7 @@ class NativeLifecycle(Protocol[RequestT, ResponseT]):
|
|||
def _lifecycle(value: object) -> NativeLifecycle[object, object] | None:
|
||||
if not callable(value):
|
||||
return None
|
||||
return cast(NativeLifecycle[object, object], value) # cast-ok: callable native entrypoint validated above
|
||||
return value # pyright: ignore[reportReturnType] # callable shape is validated by native contract tests
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -65,9 +65,7 @@ class ComponentExecution:
|
|||
|
||||
def require_supported(self) -> None:
|
||||
if self.decision is ExecutionDecision.UNSUPPORTED:
|
||||
raise RustRouteUnsupportedError(
|
||||
f"No Python or Rust implementation for {self.route_name.value}"
|
||||
)
|
||||
raise RustRouteUnsupportedError(f"No Python or Rust implementation for {self.route_name.value}")
|
||||
|
||||
def select(self, binding: NativeBinding[BindingT]) -> BindingT | None:
|
||||
self.require_supported()
|
||||
|
|
@ -75,9 +73,7 @@ class ComponentExecution:
|
|||
return None
|
||||
selected: Final = binding.load()
|
||||
if selected is None and self.decision is ExecutionDecision.RUST_REQUIRED:
|
||||
raise RustRouteUnavailableError(
|
||||
f"Rust {self.route_name.value} bridge is unavailable"
|
||||
)
|
||||
raise RustRouteUnavailableError(f"Rust {self.route_name.value} bridge is unavailable")
|
||||
return selected
|
||||
|
||||
|
||||
|
|
@ -87,10 +83,11 @@ class NativeComponent:
|
|||
capability: CapabilitySpec
|
||||
exports: tuple[str, ...]
|
||||
|
||||
def resolve(self, context: CapabilityContext = CapabilityContext()) -> ComponentExecution:
|
||||
def resolve(self, context: CapabilityContext | None = None) -> ComponentExecution:
|
||||
resolved_context: Final = context if context is not None else CapabilityContext()
|
||||
return ComponentExecution(
|
||||
route_name=self.name,
|
||||
decision=capability_decision(self.capability, context=context),
|
||||
decision=capability_decision(self.capability, context=resolved_context),
|
||||
)
|
||||
|
||||
def bind(
|
||||
|
|
|
|||
|
|
@ -28,13 +28,15 @@ class BridgeErrorContext:
|
|||
def invoke(
|
||||
*,
|
||||
native_call: Callable[[], NativeT] | None,
|
||||
python_fallback: Callable[[], ResultT],
|
||||
python_fallback: Callable[[], ResultT] | None,
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
execution: ComponentExecution,
|
||||
context: BridgeErrorContext,
|
||||
) -> ResultT:
|
||||
execution.require_supported()
|
||||
_validate_fallback(execution, python_fallback)
|
||||
if execution.decision is ExecutionDecision.PYTHON:
|
||||
assert python_fallback is not None
|
||||
return python_fallback()
|
||||
if native_call is None:
|
||||
return _unavailable_or_fallback(execution, python_fallback)
|
||||
|
|
@ -61,20 +63,22 @@ def invoke(
|
|||
async def ainvoke(
|
||||
*,
|
||||
native_call: Callable[[], Awaitable[NativeT]] | None,
|
||||
python_fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
python_fallback: Callable[[], Awaitable[ResultT]] | None,
|
||||
adapt: Callable[[NativeT], Awaitable[ResultT]],
|
||||
execution: ComponentExecution,
|
||||
context: BridgeErrorContext,
|
||||
) -> ResultT:
|
||||
execution.require_supported()
|
||||
_validate_fallback(execution, python_fallback)
|
||||
if execution.decision is ExecutionDecision.PYTHON:
|
||||
assert python_fallback is not None
|
||||
return await python_fallback()
|
||||
if native_call is None:
|
||||
return await _aunavailable_or_fallback(execution, python_fallback)
|
||||
|
||||
exceptions: Final = native_exception_types()
|
||||
if exceptions is None:
|
||||
return adapt(await native_call())
|
||||
return await adapt(await native_call())
|
||||
declined, upstream = exceptions
|
||||
unavailable: Final = native_unavailable_exception()
|
||||
host_callback: Final = native_host_callback_exception()
|
||||
|
|
@ -88,53 +92,65 @@ async def ainvoke(
|
|||
return await _adeclined_or_fallback(execution, python_fallback, error)
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
return adapt(value)
|
||||
return await adapt(value)
|
||||
|
||||
|
||||
def _unavailable_or_fallback(
|
||||
execution: ComponentExecution,
|
||||
python_fallback: Callable[[], ResultT],
|
||||
python_fallback: Callable[[], ResultT] | None,
|
||||
) -> ResultT:
|
||||
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
|
||||
assert python_fallback is not None
|
||||
return python_fallback()
|
||||
raise RustRouteUnavailableError(f"Rust {execution.route_name.value} bridge is unavailable")
|
||||
|
||||
|
||||
async def _aunavailable_or_fallback(
|
||||
execution: ComponentExecution,
|
||||
python_fallback: Callable[[], Awaitable[ResultT]],
|
||||
python_fallback: Callable[[], Awaitable[ResultT]] | None,
|
||||
) -> ResultT:
|
||||
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
|
||||
assert python_fallback is not None
|
||||
return await python_fallback()
|
||||
raise RustRouteUnavailableError(f"Rust {execution.route_name.value} bridge is unavailable")
|
||||
|
||||
|
||||
def _declined_or_fallback(
|
||||
execution: ComponentExecution,
|
||||
python_fallback: Callable[[], ResultT],
|
||||
python_fallback: Callable[[], ResultT] | None,
|
||||
error: BaseException,
|
||||
) -> ResultT:
|
||||
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
|
||||
assert python_fallback is not None
|
||||
return python_fallback()
|
||||
_raise_declined(execution, error)
|
||||
|
||||
|
||||
async def _adeclined_or_fallback(
|
||||
execution: ComponentExecution,
|
||||
python_fallback: Callable[[], Awaitable[ResultT]],
|
||||
python_fallback: Callable[[], Awaitable[ResultT]] | None,
|
||||
error: BaseException,
|
||||
) -> ResultT:
|
||||
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
|
||||
assert python_fallback is not None
|
||||
return await python_fallback()
|
||||
_raise_declined(execution, error)
|
||||
|
||||
|
||||
def _validate_fallback(execution: ComponentExecution, python_fallback: object | None) -> None:
|
||||
has_python: Final = execution.decision in (ExecutionDecision.PYTHON, ExecutionDecision.RUST_WITH_FALLBACK)
|
||||
if has_python and not callable(python_fallback):
|
||||
raise ValueError(
|
||||
f"{execution.route_name.value} declares a Python implementation but no callable fallback was supplied"
|
||||
)
|
||||
if not has_python and python_fallback is not None:
|
||||
raise ValueError(f"{execution.route_name.value} declares no Python implementation but a fallback was supplied")
|
||||
|
||||
|
||||
def _raise_declined(execution: ComponentExecution, error: BaseException) -> NoReturn:
|
||||
reason_value: Final[object] = error.args[0] if error.args else str(error)
|
||||
reason: Final = reason_value if isinstance(reason_value, str) else str(reason_value)
|
||||
raise RustRouteDeclinedError(
|
||||
f"Rust {execution.route_name.value} bridge declined the request: {reason}"
|
||||
) from error
|
||||
raise RustRouteDeclinedError(f"Rust {execution.route_name.value} bridge declined the request: {reason}") from error
|
||||
|
||||
|
||||
def _raise_host_callback(error: BaseException) -> NoReturn:
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final, cast # noqa: TID251 # native extension exposes untyped callables
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
|
@ -47,27 +49,43 @@ def _tokenizer_resource(tokenizer: str) -> str:
|
|||
raise ValueError(f"unsupported Rust tokenizer resource: {tokenizer}")
|
||||
|
||||
|
||||
async def count_input_tokens(body: bytes, tokenizer: RustTokenizer) -> InputTokenCount | None:
|
||||
async def count_input_tokens(
|
||||
*,
|
||||
body: bytes | None,
|
||||
tokenizers: Mapping[str, RustTokenizer],
|
||||
python_fallback: Callable[[], Awaitable[Mapping[str, int]]],
|
||||
) -> Mapping[str, int]:
|
||||
counter: Final = TOKEN_COUNTER.load()
|
||||
distinct_tokenizers: Final = tuple(dict.fromkeys(tokenizers.values()))
|
||||
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
async def native_call() -> Mapping[str, int]:
|
||||
assert counter is not None
|
||||
assert body is not None
|
||||
counts: Final = tuple(
|
||||
[
|
||||
_INPUT_TOKEN_COUNT.validate_python(
|
||||
await counter(
|
||||
body,
|
||||
tokenizer.kind,
|
||||
tokenizer.encoding,
|
||||
tokenizer.disabled,
|
||||
tokenizer.legacy_accounting,
|
||||
_tokenizer_resource,
|
||||
)
|
||||
).input_tokens
|
||||
for tokenizer in distinct_tokenizers
|
||||
]
|
||||
)
|
||||
counts_by_tokenizer: Final = MappingProxyType(dict(zip(distinct_tokenizers, counts)))
|
||||
return MappingProxyType({model: counts_by_tokenizer[tokenizer] for model, tokenizer in tokenizers.items()})
|
||||
|
||||
async def adapt(counts: Mapping[str, int]) -> Mapping[str, int]:
|
||||
return counts
|
||||
|
||||
return await ainvoke(
|
||||
execution=ComponentExecution(COMPONENT.name, ExecutionDecision.RUST_WITH_FALLBACK),
|
||||
native_call=(
|
||||
lambda: counter(
|
||||
body,
|
||||
tokenizer.kind,
|
||||
tokenizer.encoding,
|
||||
tokenizer.disabled,
|
||||
tokenizer.legacy_accounting,
|
||||
_tokenizer_resource,
|
||||
)
|
||||
)
|
||||
if counter is not None
|
||||
else None,
|
||||
native_call=native_call if counter is not None and body is not None else None,
|
||||
python_fallback=python_fallback,
|
||||
adapt=_INPUT_TOKEN_COUNT.validate_python,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider="", model=""),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import (
|
||||
Final,
|
||||
cast, # noqa: TID251 # native callable signatures are checked by bridge contract tests
|
||||
|
|
@ -54,7 +55,8 @@ def transcription(
|
|||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
python_fallback: Callable[[], dict[str, object]] | None,
|
||||
) -> dict[str, object]:
|
||||
execution: Final = COMPONENT.resolve(CapabilityContext(provider=custom_llm_provider or "", model=model))
|
||||
rust_transcription: Final = execution.select(_TRANSCRIPTION)
|
||||
return invoke(
|
||||
|
|
@ -73,7 +75,7 @@ def transcription(
|
|||
)
|
||||
if rust_transcription is not None
|
||||
else None,
|
||||
python_fallback=lambda: None,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda response: response,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
|
@ -89,11 +91,13 @@ async def atranscription(
|
|||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
python_fallback: Callable[[], Awaitable[dict[str, object]]] | None,
|
||||
) -> dict[str, object]:
|
||||
execution: Final = COMPONENT.resolve(CapabilityContext(provider=custom_llm_provider or "", model=model))
|
||||
rust_atranscription: Final = execution.select(_ATRANSCRIPTION)
|
||||
async def python_fallback() -> None:
|
||||
return None
|
||||
|
||||
async def adapt(response: dict[str, object]) -> dict[str, object]:
|
||||
return response
|
||||
|
||||
return await ainvoke(
|
||||
execution=execution,
|
||||
|
|
@ -112,6 +116,6 @@ async def atranscription(
|
|||
if rust_atranscription is not None
|
||||
else None,
|
||||
python_fallback=python_fallback,
|
||||
adapt=lambda response: response,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ from litellm.rust_bridge import configuration
|
|||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
rust_messages = importlib.import_module("litellm.rust_bridge.messages")
|
||||
rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader")
|
||||
|
|
@ -32,6 +31,7 @@ REQUEST_BODY: dict[str, object] = {
|
|||
"max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
PYTHON_MESSAGES_RESPONSE: dict[str, object] = {"id": "python_fallback"}
|
||||
|
||||
|
||||
class RecordingMessages:
|
||||
|
|
@ -48,6 +48,7 @@ class RecordingMessages:
|
|||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
has_agentic_hook: bool = False,
|
||||
on_request=None,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
|
|
@ -60,6 +61,8 @@ class RecordingMessages:
|
|||
"timeout_seconds": timeout_seconds,
|
||||
}
|
||||
)
|
||||
if on_request is not None:
|
||||
on_request()
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
||||
|
||||
|
|
@ -77,6 +80,7 @@ class RecordingAsyncMessages:
|
|||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
has_agentic_hook: bool = False,
|
||||
on_request=None,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
|
|
@ -89,6 +93,8 @@ class RecordingAsyncMessages:
|
|||
"timeout_seconds": timeout_seconds,
|
||||
}
|
||||
)
|
||||
if on_request is not None:
|
||||
on_request()
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
||||
|
||||
|
|
@ -110,6 +116,12 @@ class RaisingAsyncMessages:
|
|||
raise RuntimeError("upstream request failed with status 400: bad request")
|
||||
|
||||
|
||||
class DecliningAsyncMessages:
|
||||
async def __call__(self, **kwargs: object) -> dict[str, object]:
|
||||
native = pytest.importorskip("litellm.rust_bridge._native")
|
||||
raise native.RustBridgeDeclined("unsupported request")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_rust_flag():
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None)
|
||||
|
|
@ -135,7 +147,7 @@ def test_load_rust_amessages_returns_injected_impl():
|
|||
assert rust_messages.load_rust_amessages() is bridge
|
||||
|
||||
|
||||
def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
|
||||
def test_messages_wrapper_returns_fallback_when_bridge_absent(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge.bindings"),
|
||||
"get_native_bridge",
|
||||
|
|
@ -151,8 +163,10 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
|
|||
custom_llm_provider="azure_ai",
|
||||
extra_headers={},
|
||||
timeout=30.0,
|
||||
python_fallback=lambda: dict(PYTHON_MESSAGES_RESPONSE),
|
||||
adapt=lambda response: response,
|
||||
)
|
||||
assert result is None
|
||||
assert result == PYTHON_MESSAGES_RESPONSE
|
||||
|
||||
|
||||
def test_messages_wrapper_forwards_args_and_converts_timeout():
|
||||
|
|
@ -168,6 +182,8 @@ def test_messages_wrapper_forwards_args_and_converts_timeout():
|
|||
custom_llm_provider="azure_ai",
|
||||
extra_headers={"anthropic-beta": "token-efficient-tools-2025-02-19"},
|
||||
timeout=httpx.Timeout(600.0, read=42.0),
|
||||
python_fallback=lambda: pytest.fail("native request should not fall back"),
|
||||
adapt=lambda response: response,
|
||||
)
|
||||
|
||||
assert response == FAKE_MESSAGES_RESPONSE
|
||||
|
|
@ -188,6 +204,12 @@ async def test_amessages_wrapper_forwards_args():
|
|||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
async def python_fallback() -> dict[str, object]:
|
||||
pytest.fail("native request should not fall back")
|
||||
|
||||
async def adapt(response: dict[str, object]) -> dict[str, object]:
|
||||
return response
|
||||
|
||||
response = await rust_messages.amessages(
|
||||
model="claude-sonnet-4-5",
|
||||
body=REQUEST_BODY,
|
||||
|
|
@ -196,6 +218,8 @@ async def test_amessages_wrapper_forwards_args():
|
|||
custom_llm_provider="azure_ai",
|
||||
extra_headers=None,
|
||||
timeout=12.5,
|
||||
python_fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
)
|
||||
|
||||
assert response == FAKE_MESSAGES_RESPONSE
|
||||
|
|
@ -203,10 +227,9 @@ async def test_amessages_wrapper_forwards_args():
|
|||
assert bridge.calls[0]["timeout_seconds"] == 12.5
|
||||
|
||||
|
||||
def _gate(**overrides):
|
||||
async def _gate(**overrides):
|
||||
kwargs = {
|
||||
"custom_llm_provider": "azure_ai",
|
||||
"litellm_params": GenericLiteLLMParams(api_key="sk-azure"),
|
||||
"has_agentic_hook": False,
|
||||
"model": "claude-sonnet-4-5",
|
||||
"api_key": "sk-azure",
|
||||
|
|
@ -216,7 +239,28 @@ def _gate(**overrides):
|
|||
"timeout": 30.0,
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs)
|
||||
request_body = kwargs.pop("request_body")
|
||||
|
||||
async def python_fallback() -> dict[str, object]:
|
||||
return dict(PYTHON_MESSAGES_RESPONSE)
|
||||
|
||||
async def adapt(response: dict[str, object]) -> dict[str, object]:
|
||||
adapted = dict(response)
|
||||
adapted["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
|
||||
return adapted
|
||||
|
||||
return await rust_messages.amessages(
|
||||
model=kwargs["model"],
|
||||
body={key: value for key, value in request_body.items() if key != "stream"},
|
||||
has_agentic_hook=kwargs["has_agentic_hook"],
|
||||
api_key=kwargs["api_key"],
|
||||
api_base=kwargs["api_base"],
|
||||
custom_llm_provider=kwargs["custom_llm_provider"],
|
||||
extra_headers=kwargs["headers"],
|
||||
timeout=kwargs["timeout"],
|
||||
python_fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -256,9 +300,9 @@ async def test_gate_skips_rust_when_flag_absent(monkeypatch):
|
|||
bridge = ExplodingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
response = await _gate()
|
||||
|
||||
assert response is None
|
||||
assert response == PYTHON_MESSAGES_RESPONSE
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
|
|
@ -268,7 +312,7 @@ async def test_gate_uses_process_enable_without_request_override():
|
|||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
litellm.rust(True)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
response = await _gate()
|
||||
|
||||
assert response is not None
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
|
||||
|
|
@ -282,7 +326,6 @@ async def test_gate_invokes_rust_for_native_anthropic_provider():
|
|||
|
||||
response = await _gate(
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
api_key="sk-ant",
|
||||
api_base="https://api.anthropic.com",
|
||||
headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"},
|
||||
|
|
@ -302,7 +345,6 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch):
|
|||
|
||||
response = await _gate(
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
|
|
@ -317,29 +359,26 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
|
|||
|
||||
response = await _gate(
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
)
|
||||
|
||||
assert response is None
|
||||
assert response == PYTHON_MESSAGES_RESPONSE
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_unsupported_provider():
|
||||
native = pytest.importorskip("litellm.rust_bridge._native")
|
||||
async def test_gate_falls_back_for_unsupported_provider():
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=native.amessages)
|
||||
rust_messages.set_rust_messages(amessages=DecliningAsyncMessages())
|
||||
response = await _gate(custom_llm_provider="openai", api_base="http://127.0.0.1:1")
|
||||
assert response is None
|
||||
assert response == PYTHON_MESSAGES_RESPONSE
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_agentic_hook():
|
||||
native = pytest.importorskip("litellm.rust_bridge._native")
|
||||
async def test_gate_falls_back_for_agentic_hook():
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=native.amessages)
|
||||
rust_messages.set_rust_messages(amessages=DecliningAsyncMessages())
|
||||
response = await _gate(has_agentic_hook=True, api_base="http://127.0.0.1:1")
|
||||
assert response is None
|
||||
assert response == PYTHON_MESSAGES_RESPONSE
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -387,4 +426,4 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch):
|
|||
|
||||
response = await _gate()
|
||||
|
||||
assert response is None
|
||||
assert response == PYTHON_MESSAGES_RESPONSE
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ from litellm.rust_bridge.ocr import value as rust_ocr_bridge
|
|||
|
||||
def test_rust_ocr_response_retains_provider_native_response():
|
||||
provider_response = {"status": "succeeded", "analyzeResult": {"content": "native"}}
|
||||
response = rust_ocr_bridge._response(
|
||||
response = rust_ocr_bridge.adapt_response(
|
||||
{
|
||||
"pages": [],
|
||||
"model": "prebuilt-layout",
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import json
|
|||
import math
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -12,6 +13,7 @@ from litellm.caching import DualCache
|
|||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.spend_tracking import budget_reservation
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
count_request_input_tokens,
|
||||
estimate_request_max_cost,
|
||||
|
|
@ -232,58 +234,37 @@ class _FakeNative:
|
|||
RustUpstreamError = _FakeUpstream
|
||||
|
||||
|
||||
class _RecordingCounter:
|
||||
"""Stands in for one native counter; records `(tokenizer, body)` on the shared factory."""
|
||||
|
||||
def __init__(self, factory: _RecordingFactory, tokenizer: rust_token_counter.RustTokenizer) -> None:
|
||||
self.factory = factory
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
async def acount_request(self, body: bytes) -> object:
|
||||
self.factory.calls.append((self.tokenizer, body))
|
||||
return {"model": "", "input_tokens": RUST_INPUT_TOKENS_BY_TOKENIZER[self.tokenizer]}
|
||||
|
||||
|
||||
class _RecordingFactory:
|
||||
"""Stands in for the native `TokenCounter` class: called with tokenizer JSON, or `from_*_ranks`."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[tuple[rust_token_counter.RustTokenizer, bytes]] = []
|
||||
self.calls: list[tuple[str, bytes]] = []
|
||||
|
||||
def __call__(self, tokenizer_json: str) -> _RecordingCounter:
|
||||
return _RecordingCounter(self, "anthropic")
|
||||
|
||||
def from_cl100k_ranks(self, rank_file: str) -> _RecordingCounter:
|
||||
return _RecordingCounter(self, "cl100k_base")
|
||||
|
||||
def from_o200k_ranks(self, rank_file: str) -> _RecordingCounter:
|
||||
return _RecordingCounter(self, "o200k_base")
|
||||
|
||||
|
||||
class _DecliningCounter:
|
||||
async def acount_request(self, body: bytes) -> object:
|
||||
raise _FakeDeclined("unsupported content block")
|
||||
async def __call__(
|
||||
self,
|
||||
body: bytes,
|
||||
kind: str | None,
|
||||
encoding: str,
|
||||
disabled: bool,
|
||||
legacy_accounting: bool,
|
||||
resource_loader,
|
||||
) -> object:
|
||||
tokenizer: Final = kind or encoding
|
||||
self.calls.append((tokenizer, body))
|
||||
if disabled or legacy_accounting or tokenizer not in RUST_INPUT_TOKENS_BY_TOKENIZER:
|
||||
raise _FakeDeclined("unsupported tokenizer configuration")
|
||||
return {"model": "", "input_tokens": RUST_INPUT_TOKENS_BY_TOKENIZER[tokenizer]}
|
||||
|
||||
|
||||
class _DecliningFactory:
|
||||
def __call__(self, tokenizer_json: str) -> _DecliningCounter:
|
||||
return _DecliningCounter()
|
||||
|
||||
def from_cl100k_ranks(self, rank_file: str) -> _DecliningCounter:
|
||||
return _DecliningCounter()
|
||||
|
||||
def from_o200k_ranks(self, rank_file: str) -> _DecliningCounter:
|
||||
return _DecliningCounter()
|
||||
async def __call__(self, *args, **kwargs) -> object:
|
||||
raise _FakeDeclined("unsupported content block")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def rust_counter(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative())
|
||||
rust_token_counter._counter.cache_clear()
|
||||
configuration.reset_rust_configuration()
|
||||
yield
|
||||
rust_token_counter.TOKEN_COUNTER.reset()
|
||||
rust_token_counter._counter.cache_clear()
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
|
|
@ -350,7 +331,7 @@ async def test_tiktoken_o200k_models_are_counted_by_rust(rust_counter: None, mod
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multi_model_request_counts_once_per_tokenizer_and_python_for_the_rest(rust_counter: None) -> None:
|
||||
async def test_multi_model_decline_discards_every_native_count(rust_counter: None) -> None:
|
||||
factory: Final = _RecordingFactory()
|
||||
litellm.rust(True)
|
||||
rust_token_counter.TOKEN_COUNTER.override(factory)
|
||||
|
|
@ -372,16 +353,14 @@ async def test_multi_model_request_counts_once_per_tokenizer_and_python_for_the_
|
|||
request_body=body, route="/v1/chat/completions", llm_router=None, raw_body=raw_body
|
||||
)
|
||||
|
||||
assert factory.calls == [("cl100k_base", raw_body), ("anthropic", raw_body), ("o200k_base", raw_body)]
|
||||
assert dict(counts) == {
|
||||
CL100K_MODEL: RUST_INPUT_TOKENS_BY_TOKENIZER["cl100k_base"],
|
||||
"gemini/gemini-2.5-pro": RUST_INPUT_TOKENS_BY_TOKENIZER["cl100k_base"],
|
||||
ANTHROPIC_TOKENIZER_MODEL: RUST_INPUT_TOKENS,
|
||||
O200K_MODEL: RUST_INPUT_TOKENS_BY_TOKENIZER["o200k_base"],
|
||||
"gpt-5": RUST_INPUT_TOKENS_BY_TOKENIZER["o200k_base"],
|
||||
"replicate/meta/llama-2-70b-chat": python_counts["replicate/meta/llama-2-70b-chat"],
|
||||
}
|
||||
assert counts["replicate/meta/llama-2-70b-chat"] not in RUST_INPUT_TOKENS_BY_TOKENIZER.values()
|
||||
assert factory.calls == [
|
||||
("cl100k_base", raw_body),
|
||||
("anthropic", raw_body),
|
||||
("o200k_base", raw_body),
|
||||
("llama2", raw_body),
|
||||
]
|
||||
assert dict(counts) == dict(python_counts)
|
||||
assert not set(counts.values()) & set(RUST_INPUT_TOKENS_BY_TOKENIZER.values())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -431,7 +410,7 @@ async def test_direct_budget_counter_ignores_disabled_public_rollout(rust_counte
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", ("replicate/meta/llama-2-70b-chat", "meta-llama/Llama-3-8b", "text-davinci-003"))
|
||||
async def test_models_without_a_rust_tokenizer_stay_in_python(
|
||||
async def test_models_with_unsupported_native_configuration_fall_back_to_python(
|
||||
rust_counter: None, monkeypatch: pytest.MonkeyPatch, model: str
|
||||
) -> None:
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -449,6 +428,28 @@ async def test_models_without_a_rust_tokenizer_stay_in_python(
|
|||
request_body=body, route="/v1/chat/completions", llm_router=None, raw_body=json.dumps(body).encode()
|
||||
)
|
||||
|
||||
assert factory.calls == []
|
||||
tokenizer: Final = rust_token_counter.rust_tokenizer(model)
|
||||
assert factory.calls == [(tokenizer.kind or tokenizer.encoding, json.dumps(body).encode())]
|
||||
assert dict(counts) == dict(python_counts)
|
||||
assert counts[model] not in RUST_INPUT_TOKENS_BY_TOKENIZER.values()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_large_python_fallback_runs_off_event_loop(
|
||||
rust_counter: None,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
litellm.rust(False)
|
||||
body: Final = {"model": O200K_MODEL, "input": "x" * budget_reservation.TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS}
|
||||
offload: Final = AsyncMock(side_effect=lambda function, **kwargs: function(**kwargs))
|
||||
monkeypatch.setattr(budget_reservation.asyncio, "to_thread", offload)
|
||||
|
||||
counts: Final = await count_request_input_tokens(
|
||||
request_body=body,
|
||||
route="/v1/responses",
|
||||
llm_router=None,
|
||||
raw_body=json.dumps(body).encode(),
|
||||
)
|
||||
|
||||
assert counts[O200K_MODEL] > 0
|
||||
offload.assert_awaited_once()
|
||||
|
|
|
|||
|
|
@ -14,9 +14,15 @@ class _FakeNativeConnection:
|
|||
async def send_text(self, text: str) -> None:
|
||||
self.sent.append(text)
|
||||
|
||||
async def send(self, text: str) -> None:
|
||||
self.sent.append(text)
|
||||
|
||||
async def recv_text(self) -> str:
|
||||
return "response.completed"
|
||||
|
||||
async def recv(self) -> str:
|
||||
return "response.completed"
|
||||
|
||||
async def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
|
@ -39,6 +45,10 @@ class _FakeNativeBridge:
|
|||
return _FakeNativeConnection()
|
||||
|
||||
|
||||
async def _no_connection() -> _FakeNativeConnection:
|
||||
raise AssertionError("Python fallback must not connect")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_responses_websocket():
|
||||
responses_websocket.set_rust_responses_websocket(connection=None)
|
||||
|
|
@ -57,20 +67,22 @@ async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
async def test_bridge_unavailable_executes_fallback(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
configuration.rust(True)
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
|
||||
fallback = _FakeNativeConnection()
|
||||
|
||||
assert (
|
||||
await responses_websocket.connect(
|
||||
url="wss://example.test/responses",
|
||||
custom_llm_provider="openai",
|
||||
model="test",
|
||||
headers={},
|
||||
timeout=None,
|
||||
)
|
||||
is None
|
||||
)
|
||||
async def python_fallback() -> _FakeNativeConnection:
|
||||
return fallback
|
||||
|
||||
assert await responses_websocket.connect(
|
||||
url="wss://example.test/responses",
|
||||
custom_llm_provider="openai",
|
||||
model="test",
|
||||
headers={},
|
||||
timeout=None,
|
||||
python_fallback=python_fallback,
|
||||
) is fallback
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -86,6 +98,7 @@ async def test_enabled_bridge_connects_and_adapts_socket(
|
|||
model="test",
|
||||
headers={"Authorization": "Bearer key"},
|
||||
timeout=1.0,
|
||||
python_fallback=_no_connection,
|
||||
)
|
||||
|
||||
assert connection is not None
|
||||
|
|
@ -103,16 +116,19 @@ async def test_disabled_websocket_does_not_connect() -> None:
|
|||
|
||||
responses_websocket.set_rust_responses_websocket(connection=UnexpectedConnection)
|
||||
configuration.rust(False)
|
||||
assert (
|
||||
await responses_websocket.connect(
|
||||
url="ws://127.0.0.1:1",
|
||||
headers={},
|
||||
timeout=0.1,
|
||||
custom_llm_provider="openai",
|
||||
model="test",
|
||||
)
|
||||
is None
|
||||
)
|
||||
fallback = _FakeNativeConnection()
|
||||
|
||||
async def python_fallback() -> _FakeNativeConnection:
|
||||
return fallback
|
||||
|
||||
assert await responses_websocket.connect(
|
||||
url="ws://127.0.0.1:1",
|
||||
headers={},
|
||||
timeout=0.1,
|
||||
custom_llm_provider="openai",
|
||||
model="test",
|
||||
python_fallback=python_fallback,
|
||||
) is fallback
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -122,16 +138,19 @@ async def test_native_websocket_decline_falls_back_but_connection_failure_does_n
|
|||
native = pytest.importorskip("litellm.rust_bridge._native")
|
||||
responses_websocket.set_rust_responses_websocket(connection=native.ResponsesWebSocketConnection)
|
||||
configuration.rust(True)
|
||||
assert (
|
||||
await responses_websocket.connect(
|
||||
url="ws://127.0.0.1:1",
|
||||
headers={},
|
||||
timeout=0.1,
|
||||
custom_llm_provider="azure",
|
||||
model="test",
|
||||
)
|
||||
is None
|
||||
)
|
||||
fallback = _FakeNativeConnection()
|
||||
|
||||
async def python_fallback() -> _FakeNativeConnection:
|
||||
return fallback
|
||||
|
||||
assert await responses_websocket.connect(
|
||||
url="ws://127.0.0.1:1",
|
||||
headers={},
|
||||
timeout=0.1,
|
||||
custom_llm_provider="azure",
|
||||
model="test",
|
||||
python_fallback=python_fallback,
|
||||
) is fallback
|
||||
with pytest.raises(APIError):
|
||||
await responses_websocket.connect(
|
||||
url="ws://127.0.0.1:1",
|
||||
|
|
@ -139,4 +158,5 @@ async def test_native_websocket_decline_falls_back_but_connection_failure_does_n
|
|||
timeout=0.1,
|
||||
custom_llm_provider="openai",
|
||||
model="test",
|
||||
python_fallback=python_fallback,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -21,7 +21,8 @@ from litellm.rust_bridge.configuration import (
|
|||
)
|
||||
from litellm.rust_bridge.errors import RustRouteUnsupportedError
|
||||
from litellm.rust_bridge.route import NativeComponent
|
||||
from litellm.rust_bridge.token_counter import COMPONENT, definition
|
||||
from litellm.rust_bridge.token_counter import definition
|
||||
from litellm.rust_bridge.token_counter import COMPONENT
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("value", "expected"), (("1", True), ("0", False), (" 1 ", True), (" 0 ", False)))
|
||||
|
|
|
|||
|
|
@ -12,10 +12,10 @@ from litellm.rust_bridge.catalog import COMPONENTS
|
|||
from litellm.rust_bridge.configuration import (
|
||||
CapabilityContext,
|
||||
CapabilityDefinition,
|
||||
ComponentName,
|
||||
DeliveryMode,
|
||||
ExecutionDecision,
|
||||
RolloutPolicy,
|
||||
ComponentName,
|
||||
RustImplementationState,
|
||||
)
|
||||
from litellm.rust_bridge.errors import RustRouteUnavailableError, RustRouteUnsupportedError
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ import pytest
|
|||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge import bindings, runtime
|
||||
from litellm.rust_bridge.configuration import ComponentName, ExecutionDecision
|
||||
from litellm.rust_bridge.configuration import ExecutionDecision, ComponentName
|
||||
from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError, RustRouteUnsupportedError
|
||||
from litellm.rust_bridge.route import ComponentExecution
|
||||
|
||||
|
|
@ -46,7 +46,7 @@ async def _invoke(
|
|||
decision: ExecutionDecision,
|
||||
native_call: Callable[[], object] | None,
|
||||
adapt: Callable[[object], object],
|
||||
fallback: Callable[[], object] = lambda: "python",
|
||||
fallback: Callable[[], object] | None = lambda: "python",
|
||||
) -> object:
|
||||
execution: Final = ComponentExecution(route_name=ComponentName.MESSAGES, decision=decision)
|
||||
context: Final = runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model")
|
||||
|
|
@ -64,13 +64,17 @@ async def _invoke(
|
|||
return native_call()
|
||||
|
||||
async def afallback() -> object:
|
||||
assert fallback is not None
|
||||
return fallback()
|
||||
|
||||
async def aadapt(value: object) -> object:
|
||||
return adapt(value)
|
||||
|
||||
return await runtime.ainvoke(
|
||||
execution=execution,
|
||||
native_call=call if native_call is not None else None,
|
||||
python_fallback=afallback,
|
||||
adapt=adapt,
|
||||
python_fallback=afallback if fallback is not None else None,
|
||||
adapt=aadapt,
|
||||
context=context,
|
||||
)
|
||||
|
||||
|
|
@ -111,11 +115,32 @@ async def test_unavailable_and_declined_follow_policy(
|
|||
if required:
|
||||
expected: Final = RustRouteDeclinedError if isinstance(error, RustBridgeDeclined) else RustRouteUnavailableError
|
||||
with pytest.raises(expected):
|
||||
await _invoke(asynchronous, decision, native_call, str)
|
||||
await _invoke(asynchronous, decision, native_call, str, None)
|
||||
return
|
||||
assert await _invoke(asynchronous, decision, native_call, str) == "python"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("native_call", (None, lambda: (_ for _ in ()).throw(RustBridgeDeclined("declined"))))
|
||||
async def test_optional_failure_runs_python_exactly_once(
|
||||
asynchronous: bool,
|
||||
native_call: Callable[[], object] | None,
|
||||
) -> None:
|
||||
fallback_calls: Final[list[None]] = []
|
||||
|
||||
result: Final = await _invoke(
|
||||
asynchronous,
|
||||
ExecutionDecision.RUST_WITH_FALLBACK,
|
||||
native_call,
|
||||
str,
|
||||
lambda: fallback_calls.append(None) or "python",
|
||||
)
|
||||
|
||||
assert result == "python"
|
||||
assert fallback_calls == [None]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("decision", (ExecutionDecision.RUST_WITH_FALLBACK, ExecutionDecision.RUST_REQUIRED))
|
||||
|
|
@ -129,12 +154,37 @@ async def test_native_success_does_not_run_fallback(
|
|||
decision,
|
||||
lambda: value,
|
||||
lambda native: native,
|
||||
lambda: calls.append("python") or "python",
|
||||
None if decision is ExecutionDecision.RUST_REQUIRED else lambda: calls.append("python") or "python",
|
||||
)
|
||||
assert result is value
|
||||
assert calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("decision", (ExecutionDecision.PYTHON, ExecutionDecision.RUST_WITH_FALLBACK))
|
||||
async def test_python_capability_requires_fallback(asynchronous: bool, decision: ExecutionDecision) -> None:
|
||||
native_calls: Final[list[bool]] = []
|
||||
with pytest.raises(ValueError, match="declares a Python implementation"):
|
||||
await _invoke(asynchronous, decision, lambda: native_calls.append(True), str, None)
|
||||
assert native_calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_rust_required_rejects_python_fallback(asynchronous: bool) -> None:
|
||||
native_calls: Final[list[bool]] = []
|
||||
with pytest.raises(ValueError, match="declares no Python implementation"):
|
||||
await _invoke(
|
||||
asynchronous,
|
||||
ExecutionDecision.RUST_REQUIRED,
|
||||
lambda: native_calls.append(True),
|
||||
str,
|
||||
lambda: "python",
|
||||
)
|
||||
assert native_calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_upstream_failure_never_falls_back(asynchronous: bool) -> None:
|
||||
|
|
@ -170,6 +220,31 @@ async def test_adaptation_failure_never_falls_back(asynchronous: bool, error: Ex
|
|||
assert caught.value is error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_adaptation_is_awaited_and_failure_never_falls_back() -> None:
|
||||
fallback_calls: Final[list[None]] = []
|
||||
|
||||
async def adapt(_value: object) -> object:
|
||||
await asyncio.sleep(0)
|
||||
raise RuntimeError("async adapt failed")
|
||||
|
||||
execution: Final = ComponentExecution(
|
||||
route_name=ComponentName.MESSAGES,
|
||||
decision=ExecutionDecision.RUST_WITH_FALLBACK,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="async adapt failed"):
|
||||
await runtime.ainvoke(
|
||||
execution=execution,
|
||||
native_call=lambda: asyncio.sleep(0, result="native"),
|
||||
python_fallback=lambda: asyncio.sleep(0, result=fallback_calls.append(None)),
|
||||
adapt=adapt,
|
||||
context=runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model"),
|
||||
)
|
||||
|
||||
assert fallback_calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_host_callback_failure_preserves_its_cause(asynchronous: bool) -> None:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Callable, Iterator
|
||||
from collections.abc import Awaitable, Callable, Iterator, Mapping
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
|
@ -20,7 +20,6 @@ BODY: Final = json.dumps({"model": MODEL, "messages": [{"role": "user", "content
|
|||
ANTHROPIC: Final = bridge.RustTokenizer(kind="anthropic", encoding="", disabled=False, legacy_accounting=False)
|
||||
CL100K: Final = bridge.RustTokenizer(kind=None, encoding="cl100k_base", disabled=False, legacy_accounting=False)
|
||||
O200K: Final = bridge.RustTokenizer(kind=None, encoding="o200k_base", disabled=False, legacy_accounting=False)
|
||||
TOKENIZERS: Final = (ANTHROPIC, CL100K, O200K)
|
||||
|
||||
_FakeDeclined = native.RustBridgeDeclined
|
||||
_FakeUnavailable = native.RustBridgeUnavailable
|
||||
|
|
@ -34,8 +33,8 @@ class _FakeNative:
|
|||
|
||||
|
||||
class _RecordingCounter:
|
||||
def __init__(self, error: Exception | None = None) -> None:
|
||||
self.error: Final = error
|
||||
def __init__(self, errors: Mapping[str, Exception] | None = None) -> None:
|
||||
self.errors: Final = errors or {}
|
||||
self.calls: Final[list[tuple[bytes, bridge.RustTokenizer, str]]] = []
|
||||
|
||||
async def __call__(
|
||||
|
|
@ -50,10 +49,10 @@ class _RecordingCounter:
|
|||
tokenizer: Final = bridge.RustTokenizer(kind, encoding, disabled, legacy_accounting)
|
||||
resource_name: Final = kind or encoding
|
||||
self.calls.append((body, tokenizer, resource_name))
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
if error := self.errors.get(resource_name):
|
||||
raise error
|
||||
resource_loader(resource_name)
|
||||
return {"model": MODEL, "input_tokens": 42}
|
||||
return {"model": MODEL, "input_tokens": {"anthropic": 42, "cl100k_base": 17}[resource_name]}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
@ -67,61 +66,118 @@ def reset_bridge(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
|
|||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
def _fallback(result: Mapping[str, int], calls: list[None]) -> Callable[[], Awaitable[Mapping[str, int]]]:
|
||||
async def call() -> Mapping[str, int]:
|
||||
calls.append(None)
|
||||
return result
|
||||
|
||||
return call
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("tokenizer", TOKENIZERS)
|
||||
async def test_direct_bridge_bypasses_disabled_public_rollout(tokenizer: bridge.RustTokenizer) -> None:
|
||||
async def test_direct_bridge_bypasses_disabled_public_rollout() -> None:
|
||||
counter: Final = _RecordingCounter()
|
||||
bridge.TOKEN_COUNTER.override(counter)
|
||||
fallback_calls: list[None] = []
|
||||
litellm.rust(False)
|
||||
assert await bridge.count_input_tokens(BODY, tokenizer) == bridge.InputTokenCount(model=MODEL, input_tokens=42)
|
||||
assert counter.calls == [(BODY, tokenizer, tokenizer.kind or tokenizer.encoding)]
|
||||
|
||||
result: Final = await bridge.count_input_tokens(
|
||||
body=BODY,
|
||||
tokenizers={MODEL: ANTHROPIC},
|
||||
python_fallback=_fallback({MODEL: 9}, fallback_calls),
|
||||
)
|
||||
|
||||
assert result == {MODEL: 42}
|
||||
assert fallback_calls == []
|
||||
assert counter.calls == [(BODY, ANTHROPIC, "anthropic")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("tokenizer", TOKENIZERS)
|
||||
async def test_one_native_count_entrypoint_receives_configuration_and_body(tokenizer: bridge.RustTokenizer) -> None:
|
||||
async def test_native_count_deduplicates_tokenizers() -> None:
|
||||
counter: Final = _RecordingCounter()
|
||||
bridge.TOKEN_COUNTER.override(counter)
|
||||
result: Final = await bridge.count_input_tokens(BODY, tokenizer)
|
||||
assert result == bridge.InputTokenCount(model=MODEL, input_tokens=42)
|
||||
assert counter.calls == [(BODY, tokenizer, tokenizer.kind or tokenizer.encoding)]
|
||||
fallback_calls: list[None] = []
|
||||
|
||||
result: Final = await bridge.count_input_tokens(
|
||||
body=BODY,
|
||||
tokenizers={MODEL: ANTHROPIC, "other-claude": ANTHROPIC, CL100K_MODEL: CL100K},
|
||||
python_fallback=_fallback({}, fallback_calls),
|
||||
)
|
||||
|
||||
assert result == {MODEL: 42, "other-claude": 42, CL100K_MODEL: 17}
|
||||
assert fallback_calls == []
|
||||
assert counter.calls == [
|
||||
(BODY, ANTHROPIC, "anthropic"),
|
||||
(BODY, CL100K, "cl100k_base"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("error", (_FakeDeclined("unsupported"), _FakeUnavailable("resource")))
|
||||
async def test_decline_and_resource_unavailability_fall_back(error: Exception) -> None:
|
||||
bridge.TOKEN_COUNTER.override(_RecordingCounter(error))
|
||||
assert await bridge.count_input_tokens(BODY, ANTHROPIC) is None
|
||||
async def test_any_decline_or_unavailability_discards_partial_native_counts(error: Exception) -> None:
|
||||
counter: Final = _RecordingCounter({"cl100k_base": error})
|
||||
bridge.TOKEN_COUNTER.override(counter)
|
||||
fallback_calls: list[None] = []
|
||||
|
||||
result: Final = await bridge.count_input_tokens(
|
||||
body=BODY,
|
||||
tokenizers={MODEL: ANTHROPIC, CL100K_MODEL: CL100K},
|
||||
python_fallback=_fallback({MODEL: 9, CL100K_MODEL: 8}, fallback_calls),
|
||||
)
|
||||
|
||||
assert result == {MODEL: 9, CL100K_MODEL: 8}
|
||||
assert fallback_calls == [None]
|
||||
assert [call[2] for call in counter.calls] == ["anthropic", "cl100k_base"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unexpected_counting_failure_propagates() -> None:
|
||||
bridge.TOKEN_COUNTER.override(_RecordingCounter(RuntimeError("encode failed")))
|
||||
@pytest.mark.parametrize("body", (None, BODY))
|
||||
async def test_missing_body_or_binding_runs_complete_python_fallback(body: bytes | None) -> None:
|
||||
bridge.TOKEN_COUNTER.override(None)
|
||||
fallback_calls: list[None] = []
|
||||
|
||||
result: Final = await bridge.count_input_tokens(
|
||||
body=body,
|
||||
tokenizers={MODEL: ANTHROPIC},
|
||||
python_fallback=_fallback({MODEL: 9}, fallback_calls),
|
||||
)
|
||||
|
||||
assert result == {MODEL: 9}
|
||||
assert fallback_calls == [None]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unexpected_counting_failure_propagates_without_python_replay() -> None:
|
||||
counter: Final = _RecordingCounter({"anthropic": RuntimeError("encode failed")})
|
||||
bridge.TOKEN_COUNTER.override(counter)
|
||||
fallback_calls: list[None] = []
|
||||
|
||||
with pytest.raises(RuntimeError, match="encode failed"):
|
||||
await bridge.count_input_tokens(BODY, ANTHROPIC)
|
||||
await bridge.count_input_tokens(
|
||||
body=BODY,
|
||||
tokenizers={MODEL: ANTHROPIC},
|
||||
python_fallback=_fallback({MODEL: 9}, fallback_calls),
|
||||
)
|
||||
|
||||
assert fallback_calls == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "expected"),
|
||||
(
|
||||
(MODEL, ANTHROPIC),
|
||||
("gpt-4", CL100K),
|
||||
("gpt-4o", O200K),
|
||||
(CL100K_MODEL, CL100K),
|
||||
(O200K_MODEL, O200K),
|
||||
("replicate/meta/llama-2-70b-chat", bridge.RustTokenizer("llama2", "", False, False)),
|
||||
),
|
||||
)
|
||||
def test_tokenizer_configuration_matches_python_selection(
|
||||
model: str, expected: bridge.RustTokenizer
|
||||
) -> None:
|
||||
bridge.TOKEN_COUNTER.override(_RecordingCounter())
|
||||
def test_tokenizer_configuration_matches_python_selection(model: str, expected: bridge.RustTokenizer) -> None:
|
||||
assert bridge.rust_tokenizer(model) == expected
|
||||
|
||||
|
||||
def test_unsupported_configuration_is_left_for_native_admission(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
bridge.TOKEN_COUNTER.override(_RecordingCounter())
|
||||
monkeypatch.setattr(litellm, "disable_token_counter", True)
|
||||
tokenizer: Final = bridge.rust_tokenizer(MODEL)
|
||||
assert tokenizer == bridge.RustTokenizer("anthropic", "", True, False)
|
||||
assert bridge.rust_tokenizer(MODEL) == bridge.RustTokenizer("anthropic", "", True, False)
|
||||
|
||||
|
||||
PARITY_REQUESTS: Final[tuple[dict[str, object], ...]] = (
|
||||
|
|
@ -149,17 +205,31 @@ async def test_native_count_matches_python_budget_counter(
|
|||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
|
||||
bridge.TOKEN_COUNTER.reset()
|
||||
body: Final = json.dumps(request_body).replace(MODEL, model)
|
||||
rust_count: Final = await bridge.count_input_tokens(body.encode(), tokenizer)
|
||||
python_count: Final = _count_input_tokens(request_body=json.loads(body), model=model)
|
||||
assert rust_count is not None
|
||||
assert rust_count.input_tokens == python_count
|
||||
|
||||
result: Final = await bridge.count_input_tokens(
|
||||
body=body.encode(),
|
||||
tokenizers={model: tokenizer},
|
||||
python_fallback=_fallback({}, []),
|
||||
)
|
||||
|
||||
assert result == {model: python_count}
|
||||
|
||||
|
||||
@pytest.mark.requires_rust_extension
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_declines_unsupported_request_without_loading_resources(
|
||||
async def test_native_declines_unsupported_request_and_runs_complete_fallback(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
|
||||
bridge.TOKEN_COUNTER.reset()
|
||||
assert await bridge.count_input_tokens(b'{"input":1.5}', ANTHROPIC) is None
|
||||
fallback_calls: list[None] = []
|
||||
|
||||
result: Final = await bridge.count_input_tokens(
|
||||
body=b'{"input":1.5}',
|
||||
tokenizers={MODEL: ANTHROPIC},
|
||||
python_fallback=_fallback({MODEL: 7}, fallback_calls),
|
||||
)
|
||||
|
||||
assert result == {MODEL: 7}
|
||||
assert fallback_calls == [None]
|
||||
|
|
|
|||
|
|
@ -71,6 +71,7 @@ def test_enabled_sync_bridge_receives_audio(enabled: bool) -> None:
|
|||
extra_headers=None,
|
||||
optional_params={"temperature": 0},
|
||||
timeout=5.0,
|
||||
python_fallback=None,
|
||||
)
|
||||
assert result == {"text": "hello"}
|
||||
assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"}
|
||||
|
|
@ -90,6 +91,7 @@ async def test_enabled_async_bridge(enabled: bool) -> None:
|
|||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
python_fallback=None,
|
||||
)
|
||||
assert result == {"text": "async"}
|
||||
|
||||
|
|
@ -214,6 +216,7 @@ async def test_bedrock_transcription_errors_never_fall_back(
|
|||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
python_fallback=None,
|
||||
)
|
||||
with pytest.raises(expected, match=message):
|
||||
await rust_bridge.atranscription(
|
||||
|
|
@ -225,6 +228,7 @@ async def test_bedrock_transcription_errors_never_fall_back(
|
|||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
python_fallback=None,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -239,29 +243,29 @@ async def test_python_transcription_skips_rust_when_enabled() -> None:
|
|||
pytest.fail("Python provider must not call Rust")
|
||||
|
||||
rust_bridge.configure_rust_transcription(transcription=unexpected, atranscription=aunexpected)
|
||||
assert (
|
||||
rust_bridge.transcription(
|
||||
model="model",
|
||||
audio={},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider="openai",
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert (
|
||||
await rust_bridge.atranscription(
|
||||
model="model",
|
||||
audio={},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider="openai",
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert rust_bridge.transcription(
|
||||
model="model",
|
||||
audio={},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider="openai",
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
python_fallback=lambda: {"text": "python"},
|
||||
) == {"text": "python"}
|
||||
|
||||
async def python_fallback() -> dict[str, object]:
|
||||
return {"text": "python"}
|
||||
|
||||
assert await rust_bridge.atranscription(
|
||||
model="model",
|
||||
audio={},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider="openai",
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
python_fallback=python_fallback,
|
||||
) == {"text": "python"}
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from typing import Final
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
from litellm.rust_bridge import _native
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
|
||||
|
|
@ -79,15 +79,15 @@ def test_native_ocr_with_compiled_rust_extension(
|
|||
host: Final = str(address[0])
|
||||
port: Final = int(address[1])
|
||||
|
||||
response: Final = rust_ocr_bridge.ocr(
|
||||
model="mistral-ocr-latest",
|
||||
document={"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
|
||||
response: Final = _native.ocr(
|
||||
"mistral-ocr-latest",
|
||||
{"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
|
||||
api_key="test-key",
|
||||
api_base=f"http://{host}:{port}",
|
||||
custom_llm_provider="mistral",
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
timeout_seconds=None,
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from litellm.rust_bridge import _native
|
|||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.catalog import NATIVE_EXPORTS
|
||||
from litellm.rust_bridge.chat_completions.lifecycle import LIFECYCLE as CHAT_COMPLETIONS
|
||||
from litellm.rust_bridge.configuration import ComponentName, ExecutionDecision
|
||||
from litellm.rust_bridge.configuration import ExecutionDecision, ComponentName
|
||||
from litellm.rust_bridge.embeddings.lifecycle import LIFECYCLE as EMBEDDINGS
|
||||
from litellm.rust_bridge.image_edit.lifecycle import LIFECYCLE as IMAGE_EDIT
|
||||
from litellm.rust_bridge.image_generation.lifecycle import LIFECYCLE as IMAGE_GENERATION
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue