mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
fix(rust): satisfy CI quality gates
This commit is contained in:
parent
6f8f07433d
commit
b02c53cebc
20 changed files with 194 additions and 112 deletions
2
Makefile
2
Makefile
|
|
@ -332,7 +332,7 @@ test-rust-python: install-rust-python-test-deps
|
|||
PYTHONPATH="$(CURDIR):$$site_packages$${PYTHONPATH:+:$$PYTHONPATH}" \
|
||||
LITELLM_LOCAL_MODEL_COST_MAP=True \
|
||||
cargo test --manifest-path litellm-rust/Cargo.toml \
|
||||
-p litellm-python-interop --tests --locked -- --include-ignored
|
||||
-p litellm-python-interop -p litellm-python-bridge --tests --locked -- --include-ignored
|
||||
|
||||
lint-rust-python-fixtures:
|
||||
$(UV) tool run --from ruff==0.15.3 ruff check --config ruff-tests.toml litellm-rust/crates/python-interop/tests
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ async fn handle(
|
|||
.map_err(MessagesRouteError::from)?
|
||||
{
|
||||
service::MessagesResponse::Json(body) => Ok(Json(body).into_response()),
|
||||
service::MessagesResponse::Stream(upstream) => stream_response(upstream),
|
||||
service::MessagesResponse::Stream(upstream) => stream_response(*upstream),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ impl TerminalDispatcher for GatewayTerminalDispatcher {
|
|||
|
||||
pub(crate) enum MessagesResponse {
|
||||
Json(Value),
|
||||
Stream(StreamingCall),
|
||||
Stream(Box<StreamingCall>),
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
|
|
@ -129,6 +129,7 @@ pub async fn run(
|
|||
if request.body.get("stream").and_then(Value::as_bool) == Some(true) {
|
||||
return lifecycle::messages_stream(services, request, Options::default(), context)
|
||||
.await
|
||||
.map(Box::new)
|
||||
.map(MessagesResponse::Stream);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -114,11 +114,13 @@ async fn bridge(
|
|||
let _ = service::run(
|
||||
&router,
|
||||
&pool,
|
||||
&model,
|
||||
None,
|
||||
loggers,
|
||||
new_call_id(),
|
||||
metadata,
|
||||
service::RealtimeCall {
|
||||
model,
|
||||
idle_timeout: None,
|
||||
loggers,
|
||||
call_id: new_call_id(),
|
||||
metadata,
|
||||
},
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -26,6 +26,14 @@ struct GatewayRealtimeServices {
|
|||
runner: CustomLoggerRunner,
|
||||
}
|
||||
|
||||
pub struct RealtimeCall {
|
||||
pub model: String,
|
||||
pub idle_timeout: Option<Duration>,
|
||||
pub loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
|
||||
pub call_id: String,
|
||||
pub metadata: RequestMetadata,
|
||||
}
|
||||
|
||||
impl Clock for GatewayRealtimeServices {
|
||||
fn now(&self) -> f64 {
|
||||
std::time::SystemTime::now()
|
||||
|
|
@ -49,11 +57,7 @@ impl TerminalDispatcher for GatewayRealtimeServices {
|
|||
pub async fn run<In, Out>(
|
||||
router: &Router,
|
||||
pool: &RealtimePool,
|
||||
model: &str,
|
||||
idle_timeout: Option<Duration>,
|
||||
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
|
||||
call_id: String,
|
||||
metadata: RequestMetadata,
|
||||
call: RealtimeCall,
|
||||
client_in: In,
|
||||
client_out: Out,
|
||||
) -> Result<ExecutedCall<(), Error>, Error>
|
||||
|
|
@ -62,8 +66,15 @@ where
|
|||
Out: Sink<RealtimeEvent> + Unpin + Send,
|
||||
<Out as Sink<RealtimeEvent>>::Error: std::fmt::Display,
|
||||
{
|
||||
let RealtimeCall {
|
||||
model,
|
||||
idle_timeout,
|
||||
loggers,
|
||||
call_id,
|
||||
metadata,
|
||||
} = call;
|
||||
let deployment = router
|
||||
.get_available_deployment(model)
|
||||
.get_available_deployment(&model)
|
||||
.ok_or_else(|| Error::Routing(format!("no deployment available for model '{model}'")))?;
|
||||
let params = &deployment.litellm_params;
|
||||
let connection = upstream_key(
|
||||
|
|
|
|||
|
|
@ -190,8 +190,7 @@ impl CallLifecycle {
|
|||
ProviderFuture: Future<Output = Result<Resp, Error>>,
|
||||
{
|
||||
self.run_with_usage(
|
||||
context,
|
||||
request,
|
||||
(context, request),
|
||||
policy,
|
||||
dispatcher,
|
||||
clock,
|
||||
|
|
@ -213,8 +212,7 @@ impl CallLifecycle {
|
|||
ResponseUsage,
|
||||
>(
|
||||
&self,
|
||||
mut context: CallLifecycleContext,
|
||||
request: InitialReq,
|
||||
input: (CallLifecycleContext, InitialReq),
|
||||
policy: &Policy,
|
||||
dispatcher: &Dispatcher,
|
||||
clock: &ClockImpl,
|
||||
|
|
@ -230,6 +228,7 @@ impl CallLifecycle {
|
|||
ProviderFuture: Future<Output = Result<Resp, Error>>,
|
||||
ResponseUsage: FnOnce(&Resp) -> Option<Usage>,
|
||||
{
|
||||
let (mut context, request) = input;
|
||||
let start_time = clock.now();
|
||||
let request = match policy.async_pre_call_hook(&context, request).await {
|
||||
ActionResult::Continue(request) | ActionResult::Replace(request) => request,
|
||||
|
|
|
|||
|
|
@ -119,7 +119,7 @@ impl StreamingCompletion {
|
|||
impl Drop for StreamingCompletion {
|
||||
fn drop(&mut self) {
|
||||
if self.receiver.is_some() {
|
||||
let _ = self.spawn();
|
||||
drop(self.spawn());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -141,8 +141,7 @@ pub async fn messages<S: MessagesServices>(
|
|||
) -> ExecutedCall<AnthropicMessagesResponse, Error> {
|
||||
CallLifecycle
|
||||
.run_with_usage(
|
||||
context,
|
||||
request,
|
||||
(context, request),
|
||||
services,
|
||||
services,
|
||||
services,
|
||||
|
|
|
|||
|
|
@ -254,7 +254,7 @@ where
|
|||
for outbound in OPENAI_REALTIME_CONFIG.transform_realtime_request(&event, model)?.events {
|
||||
let payload = serde_json::to_string(&outbound)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
|
||||
upstream_tx.send(Message::Text(payload.into())).await.map_err(ws_transport_error)?;
|
||||
upstream_tx.send(Message::Text(payload)).await.map_err(ws_transport_error)?;
|
||||
}
|
||||
observation.observe_client(&event);
|
||||
}
|
||||
|
|
@ -635,16 +635,12 @@ mod tests {
|
|||
let events = events.clone();
|
||||
tokio::spawn(async move {
|
||||
let mut socket = accept_async(stream).await.unwrap();
|
||||
socket.send(Message::Text(json!({"type":"session.created","session":{"id":"sess-core","model":"upstream-model"}}).to_string().into())).await.unwrap();
|
||||
socket.send(Message::Text(json!({"type":"session.created","session":{"id":"sess-core","model":"upstream-model"}}).to_string())).await.unwrap();
|
||||
while let Some(Ok(Message::Text(text))) = socket.next().await {
|
||||
let event: RealtimeEvent = serde_json::from_str(&text).unwrap();
|
||||
if event.event_type == "response.create" {
|
||||
for event in &events {
|
||||
if socket
|
||||
.send(Message::Text(event.to_string().into()))
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
if socket.send(Message::Text(event.to_string())).await.is_err() {
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -317,7 +317,7 @@ async fn send_provider_event(
|
|||
let payload = serde_json::to_string(&outbound)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
|
||||
upstream
|
||||
.send(Message::Text(payload.into()))
|
||||
.send(Message::Text(payload))
|
||||
.await
|
||||
.map_err(ws_transport_error)?;
|
||||
}
|
||||
|
|
@ -548,10 +548,7 @@ mod tests {
|
|||
let (stream, _) = listener.accept().await.unwrap();
|
||||
let mut socket = accept_async(stream).await.unwrap();
|
||||
for frame in frames {
|
||||
socket
|
||||
.send(Message::Text(frame.to_string().into()))
|
||||
.await
|
||||
.unwrap();
|
||||
socket.send(Message::Text(frame.to_string())).await.unwrap();
|
||||
}
|
||||
if remain_open {
|
||||
futures_util::future::pending::<()>().await;
|
||||
|
|
@ -562,10 +559,12 @@ mod tests {
|
|||
(format!("http://{address}"), task)
|
||||
}
|
||||
|
||||
fn input() -> (
|
||||
type InputChannel = (
|
||||
futures_channel::mpsc::UnboundedSender<Result<ResponsesWsEvent, Error>>,
|
||||
futures_channel::mpsc::UnboundedReceiver<Result<ResponsesWsEvent, Error>>,
|
||||
) {
|
||||
);
|
||||
|
||||
fn input() -> InputChannel {
|
||||
futures_channel::mpsc::unbounded()
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -78,6 +78,12 @@ pub(crate) fn messages_provider_error_to_pyerr(err: Error) -> PyErr {
|
|||
}
|
||||
}
|
||||
|
||||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
let py = module.py();
|
||||
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
|
||||
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
@ -120,9 +126,3 @@ mod tests {
|
|||
});
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
let py = module.py();
|
||||
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
|
||||
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -902,6 +902,7 @@ mod tests {
|
|||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires the Python SDK and its dependencies on PYTHONPATH"]
|
||||
fn callback_decline_is_terminal_and_identity_is_reused() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -984,6 +985,7 @@ asyncio.run(exercise())
|
|||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires the Python SDK and its dependencies on PYTHONPATH"]
|
||||
fn native_send_owns_state_without_the_python_driver() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -1081,6 +1083,7 @@ asyncio.run(exercise())
|
|||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires the Python SDK and its dependencies on PYTHONPATH"]
|
||||
fn retains_identity_independent_wire_roots_and_collects_cycles() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -1170,6 +1173,7 @@ assert alive() is None
|
|||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires the Python SDK and its dependencies on PYTHONPATH"]
|
||||
fn async_callbacks_are_inline_and_unsupported_requests_never_call_them() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ the LLM doesn't make a tool call, and we need to return a stream to the user.
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Callable, Mapping
|
||||
from collections.abc import AsyncIterator, Callable, Iterator, Mapping
|
||||
from types import TracebackType
|
||||
from typing import Any, Final, cast
|
||||
|
||||
|
|
@ -193,10 +193,10 @@ class FakeAnthropicMessagesStreamIterator:
|
|||
|
||||
return chunks
|
||||
|
||||
def __aiter__(self):
|
||||
def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
async def __anext__(self) -> bytes:
|
||||
if self.current_index >= len(self.chunks):
|
||||
self.close()
|
||||
raise StopAsyncIteration
|
||||
|
|
@ -205,10 +205,10 @@ class FakeAnthropicMessagesStreamIterator:
|
|||
self.current_index += 1
|
||||
return chunk
|
||||
|
||||
def __iter__(self):
|
||||
def __iter__(self) -> Iterator[bytes]:
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
def __next__(self) -> bytes:
|
||||
if self.current_index >= len(self.chunks):
|
||||
self.close()
|
||||
raise StopIteration
|
||||
|
|
|
|||
|
|
@ -2456,11 +2456,11 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
|
||||
from litellm.rust_bridge.bindings import native_exception_types
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, _raise_upstream
|
||||
from litellm.rust_bridge.runtime import BridgeErrorContext, raise_upstream
|
||||
|
||||
exception_types: Final = native_exception_types()
|
||||
if exception_types is not None and isinstance(rust_error, exception_types[1]):
|
||||
_raise_upstream(
|
||||
raise_upstream(
|
||||
rust_error,
|
||||
BridgeErrorContext(route="messages", provider=custom_llm_provider, model=model),
|
||||
)
|
||||
|
|
@ -2472,9 +2472,8 @@ class BaseLLMHTTPHandler:
|
|||
if rust_response is None:
|
||||
return None
|
||||
|
||||
response_obj: Final = cast(AnthropicMessagesResponse, rust_response)
|
||||
response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
|
||||
return response_obj
|
||||
rust_response["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
|
||||
return rust_response
|
||||
|
||||
@staticmethod
|
||||
def _rust_anthropic_messages_fake_stream(
|
||||
|
|
@ -2488,12 +2487,9 @@ class BaseLLMHTTPHandler:
|
|||
AnthropicMessagesStreamingResponse,
|
||||
)
|
||||
|
||||
completion_stream = cast(
|
||||
AsyncIterator[bytes],
|
||||
FakeAnthropicMessagesStreamIterator(
|
||||
response=rust_response,
|
||||
on_complete=getattr(rust_response, "complete", None),
|
||||
),
|
||||
completion_stream: Final[AsyncIterator[bytes]] = FakeAnthropicMessagesStreamIterator(
|
||||
response=rust_response,
|
||||
on_complete=getattr(rust_response, "complete", None),
|
||||
)
|
||||
hidden_params: Final = AnthropicMessagesStreamHiddenParams(additional_headers={"x-litellm-rust": "true"})
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
|
|
@ -6510,7 +6506,6 @@ class BaseLLMHTTPHandler:
|
|||
return
|
||||
|
||||
import websockets
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
litellm_params: Final = GenericLiteLLMParams.model_validate(
|
||||
{
|
||||
|
|
@ -6610,7 +6605,7 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
streaming: Final = ResponsesWebSocketStreaming(
|
||||
websocket=websocket,
|
||||
backend_ws=cast(ClientConnection, backend_ws),
|
||||
backend_ws=backend_ws,
|
||||
logging_obj=logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=_request_data,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
|
|
@ -13,7 +14,7 @@ def invoke_terminal(
|
|||
action: str,
|
||||
roots: object,
|
||||
logger: object,
|
||||
record: dict[str, object] | None,
|
||||
record: Mapping[str, object] | None,
|
||||
value: object,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
|
|
|
|||
|
|
@ -16,7 +16,12 @@ import inspect
|
|||
import json
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Final,
|
||||
Protocol,
|
||||
cast, # noqa: TID251 # native callables require runtime signature narrowing
|
||||
)
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
|
@ -27,7 +32,6 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
convert_to_model_response_object,
|
||||
)
|
||||
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
|
||||
from litellm.rust_bridge._lifecycle import invoke_terminal
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.loader import get_native_bridge
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
|
@ -57,6 +61,36 @@ class RustAchatCompletions(Protocol):
|
|||
raise NotImplementedError
|
||||
|
||||
|
||||
class LegacyRustChatCompletions(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> Mapping[str, object]: ...
|
||||
|
||||
|
||||
class LegacyRustAchatCompletions(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> Awaitable[Mapping[str, object]]: ...
|
||||
|
||||
|
||||
class RustChatCompletionsDecline(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
|
|
@ -108,8 +142,8 @@ _UNSET: Final[_Unset] = _Unset()
|
|||
|
||||
@dataclass(slots=True)
|
||||
class _RustChatCompletionsState:
|
||||
chat_completions: RustChatCompletions | None = None
|
||||
achat_completions: RustAchatCompletions | None = None
|
||||
chat_completions: RustChatCompletions | LegacyRustChatCompletions | None = None
|
||||
achat_completions: RustAchatCompletions | LegacyRustAchatCompletions | None = None
|
||||
decline: RustChatCompletionsDecline | None = None
|
||||
|
||||
|
||||
|
|
@ -118,8 +152,8 @@ _STATE: Final[_RustChatCompletionsState] = _RustChatCompletionsState()
|
|||
|
||||
def set_rust_chat_completions(
|
||||
*,
|
||||
chat_completions: RustChatCompletions | None | _Unset = _UNSET,
|
||||
achat_completions: RustAchatCompletions | None | _Unset = _UNSET,
|
||||
chat_completions: RustChatCompletions | LegacyRustChatCompletions | None | _Unset = _UNSET,
|
||||
achat_completions: RustAchatCompletions | LegacyRustAchatCompletions | None | _Unset = _UNSET,
|
||||
decline: RustChatCompletionsDecline | None | _Unset = _UNSET,
|
||||
) -> None:
|
||||
"""Inject the native callables, so tests can supply a double instead of
|
||||
|
|
@ -132,7 +166,7 @@ def set_rust_chat_completions(
|
|||
_STATE.decline = decline
|
||||
|
||||
|
||||
def load_rust_chat_completions() -> RustChatCompletions | None:
|
||||
def load_rust_chat_completions() -> RustChatCompletions | LegacyRustChatCompletions | None:
|
||||
if _STATE.chat_completions is not None:
|
||||
return _STATE.chat_completions
|
||||
native_bridge: Final = get_native_bridge()
|
||||
|
|
@ -142,7 +176,7 @@ def load_rust_chat_completions() -> RustChatCompletions | None:
|
|||
return loaded
|
||||
|
||||
|
||||
def load_rust_achat_completions() -> RustAchatCompletions | None:
|
||||
def load_rust_achat_completions() -> RustAchatCompletions | LegacyRustAchatCompletions | None:
|
||||
if _STATE.achat_completions is not None:
|
||||
return _STATE.achat_completions
|
||||
native_bridge: Final = get_native_bridge()
|
||||
|
|
@ -332,7 +366,10 @@ def chat_completions(
|
|||
return None
|
||||
try:
|
||||
if _STATE.chat_completions is not None and _uses_argument_bag(rust_chat_completions):
|
||||
rust_result: Final = rust_chat_completions(
|
||||
argument_bag_call: Final = cast( # cast-ok: signature inspection selected the argument-bag callable
|
||||
RustChatCompletions, rust_chat_completions
|
||||
)
|
||||
rust_result: Final = argument_bag_call(
|
||||
_arguments(
|
||||
arguments,
|
||||
model,
|
||||
|
|
@ -349,7 +386,10 @@ def chat_completions(
|
|||
)
|
||||
return rust_result
|
||||
if _STATE.chat_completions is not None:
|
||||
rust_response: Final = rust_chat_completions(
|
||||
legacy: Final = cast( # cast-ok: signature inspection selected the legacy injected callable
|
||||
LegacyRustChatCompletions, rust_chat_completions
|
||||
)
|
||||
rust_response: Final = legacy(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -362,7 +402,8 @@ def chat_completions(
|
|||
if on_response is not None:
|
||||
on_response(rust_response)
|
||||
return build_model_response(rust_response, model_response)
|
||||
return rust_chat_completions(
|
||||
native_call: Final = cast(RustChatCompletions, rust_chat_completions) # cast-ok: native ABI uses argument bag
|
||||
return native_call(
|
||||
_arguments(
|
||||
arguments,
|
||||
model,
|
||||
|
|
@ -403,7 +444,10 @@ async def achat_completions(
|
|||
return None
|
||||
try:
|
||||
if _STATE.achat_completions is not None and _uses_argument_bag(rust_achat_completions):
|
||||
rust_result: Final = await rust_achat_completions(
|
||||
argument_bag_call: Final = cast( # cast-ok: signature inspection selected the argument-bag callable
|
||||
RustAchatCompletions, rust_achat_completions
|
||||
)
|
||||
rust_result: Final = await argument_bag_call(
|
||||
_arguments(
|
||||
arguments,
|
||||
model,
|
||||
|
|
@ -420,7 +464,10 @@ async def achat_completions(
|
|||
)
|
||||
return rust_result
|
||||
if _STATE.achat_completions is not None:
|
||||
rust_response: Final = await rust_achat_completions(
|
||||
legacy: Final = cast( # cast-ok: signature inspection selected the legacy injected callable
|
||||
LegacyRustAchatCompletions, rust_achat_completions
|
||||
)
|
||||
rust_response: Final = await legacy(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -433,7 +480,10 @@ async def achat_completions(
|
|||
if on_response is not None:
|
||||
on_response(rust_response)
|
||||
return build_model_response(rust_response, model_response)
|
||||
return await rust_achat_completions(
|
||||
native_call: Final = cast( # cast-ok: native ABI uses argument bag
|
||||
RustAchatCompletions, rust_achat_completions
|
||||
)
|
||||
return await native_call(
|
||||
_arguments(
|
||||
arguments,
|
||||
model,
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ from __future__ import annotations
|
|||
|
||||
from collections.abc import Awaitable
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final, Protocol, cast
|
||||
|
||||
import httpx
|
||||
|
|
@ -11,15 +11,19 @@ from litellm.rust_bridge._lifecycle import (
|
|||
initialize_logging as initialize_lifecycle_logging,
|
||||
)
|
||||
from litellm.rust_bridge._lifecycle import invoke_terminal
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
|
||||
|
||||
class RustMessages(Protocol):
|
||||
def __call__(self, arguments: dict[str, object]) -> dict[str, object]: ...
|
||||
def __call__(self, arguments: dict[str, object]) -> AnthropicMessagesResponse: ...
|
||||
|
||||
|
||||
class RustAmessages(Protocol):
|
||||
def __call__(self, arguments: dict[str, object]) -> Awaitable[dict[str, object]]: ...
|
||||
def __call__(self, arguments: dict[str, object]) -> Awaitable[AnthropicMessagesResponse]: ...
|
||||
|
||||
|
||||
class _MessagesLogging(Protocol):
|
||||
|
|
@ -44,6 +48,18 @@ class _RustMessagesState:
|
|||
_STATE: Final = _RustMessagesState()
|
||||
|
||||
|
||||
def _as_messages(value: object) -> RustMessages | None:
|
||||
return cast(RustMessages, value) if callable(value) else None # cast-ok: callable native binding
|
||||
|
||||
|
||||
def _as_amessages(value: object) -> RustAmessages | None:
|
||||
return cast(RustAmessages, value) if callable(value) else None # cast-ok: callable native binding
|
||||
|
||||
|
||||
_MESSAGES: Final = NativeBinding("messages", validate=_as_messages)
|
||||
_AMESSAGES: Final = NativeBinding("amessages", validate=_as_amessages)
|
||||
|
||||
|
||||
def set_rust_messages(
|
||||
*,
|
||||
messages: RustMessages | None | _Unset = _UNSET,
|
||||
|
|
@ -58,19 +74,13 @@ def set_rust_messages(
|
|||
def load_rust_messages() -> RustMessages | None:
|
||||
if _STATE.messages is not None:
|
||||
return _STATE.messages
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
bridge: Final = get_native_bridge()
|
||||
return cast(RustMessages, getattr(bridge, "messages", None)) if bridge is not None else None
|
||||
return _MESSAGES.load()
|
||||
|
||||
|
||||
def load_rust_amessages() -> RustAmessages | None:
|
||||
if _STATE.amessages is not None:
|
||||
return _STATE.amessages
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
bridge: Final = get_native_bridge()
|
||||
return cast(RustAmessages, getattr(bridge, "amessages", None)) if bridge is not None else None
|
||||
return _AMESSAGES.load()
|
||||
|
||||
|
||||
def initialize_logging(arguments: dict[str, object], asynchronous: bool) -> object:
|
||||
|
|
@ -78,10 +88,12 @@ def initialize_logging(arguments: dict[str, object], asynchronous: bool) -> obje
|
|||
|
||||
|
||||
class _RetainedMessagesResponse(dict[str, object]):
|
||||
def __init__(self, response: dict[str, object], roots: object, logger: object, start_time: datetime) -> None:
|
||||
def __init__(
|
||||
self, response: AnthropicMessagesResponse, roots: object, logger: _MessagesLogging, start_time: datetime
|
||||
) -> None:
|
||||
super().__init__(response)
|
||||
self._roots = roots
|
||||
self._logger = cast(_MessagesLogging, logger)
|
||||
self._logger = logger
|
||||
self._start_time = start_time
|
||||
self._completed = False
|
||||
|
||||
|
|
@ -95,7 +107,7 @@ class _RetainedMessagesResponse(dict[str, object]):
|
|||
self
|
||||
)
|
||||
self._logger.model_call_details["complete_streaming_response"] = complete_response
|
||||
end_time = datetime.now()
|
||||
end_time = datetime.now(tz=self._start_time.tzinfo or timezone.utc)
|
||||
try:
|
||||
invoke_terminal(
|
||||
"async_success",
|
||||
|
|
@ -123,9 +135,11 @@ class _RetainedMessagesResponse(dict[str, object]):
|
|||
|
||||
|
||||
def retain_stream_response(
|
||||
response: dict[str, object], roots: object, logger: object, start_time: datetime
|
||||
) -> dict[str, object]:
|
||||
return _RetainedMessagesResponse(response, roots, logger, start_time)
|
||||
response: AnthropicMessagesResponse, roots: object, logger: _MessagesLogging, start_time: datetime
|
||||
) -> AnthropicMessagesResponse:
|
||||
return cast( # cast-ok: dict subclass preserves the Anthropic response mapping contract
|
||||
AnthropicMessagesResponse, _RetainedMessagesResponse(response, roots, logger, start_time)
|
||||
)
|
||||
|
||||
|
||||
def _arguments(
|
||||
|
|
@ -138,7 +152,7 @@ def _arguments(
|
|||
extra_headers: dict[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
return { # mutable-ok: the native bridge requires a concrete argument bag
|
||||
**arguments,
|
||||
"model": model,
|
||||
"body": body,
|
||||
|
|
@ -160,7 +174,7 @@ def messages(
|
|||
extra_headers: dict[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
arguments: dict[str, object] | None = None,
|
||||
) -> dict[str, object] | None:
|
||||
) -> AnthropicMessagesResponse | None:
|
||||
implementation: Final = load_rust_messages()
|
||||
if implementation is None:
|
||||
return None
|
||||
|
|
@ -181,7 +195,7 @@ async def amessages(
|
|||
extra_headers: dict[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
arguments: dict[str, object] | None = None,
|
||||
) -> dict[str, object] | None:
|
||||
) -> AnthropicMessagesResponse | None:
|
||||
implementation: Final = load_rust_amessages()
|
||||
if implementation is None:
|
||||
return None
|
||||
|
|
@ -192,7 +206,7 @@ async def amessages(
|
|||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
__all__ = (
|
||||
"amessages",
|
||||
"initialize_logging",
|
||||
"invoke_terminal",
|
||||
|
|
@ -201,4 +215,4 @@ __all__ = [
|
|||
"messages",
|
||||
"retain_stream_response",
|
||||
"set_rust_messages",
|
||||
]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import traceback
|
||||
from collections.abc import Awaitable
|
||||
from collections.abc import Awaitable, Mapping
|
||||
from contextvars import copy_context
|
||||
from datetime import datetime
|
||||
from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables
|
||||
|
|
@ -80,7 +80,7 @@ def initialize_logging(arguments: dict[str, object], asynchronous: bool, route:
|
|||
(
|
||||
*callbacks,
|
||||
*cast( # cast-ok: per-call callback list is a legacy untyped boundary
|
||||
list, arguments.get("success_callback") or []
|
||||
list, arguments.get("success_callback") or ()
|
||||
), # cast-ok: per-call callback list is a legacy untyped boundary
|
||||
) # cast-ok: per-call callback list is a legacy untyped boundary
|
||||
) # cast-ok: per-call callback list is a legacy untyped boundary # mutable-ok: deduplication uses dict keys
|
||||
|
|
@ -90,7 +90,7 @@ def initialize_logging(arguments: dict[str, object], asynchronous: bool, route:
|
|||
(
|
||||
*callbacks,
|
||||
*cast( # cast-ok: per-call callback list is a legacy untyped boundary
|
||||
list, arguments.get("failure_callback") or []
|
||||
list, arguments.get("failure_callback") or ()
|
||||
), # cast-ok: per-call callback list is a legacy untyped boundary
|
||||
) # cast-ok: per-call callback list is a legacy untyped boundary
|
||||
) # cast-ok: per-call callback list is a legacy untyped boundary # mutable-ok: deduplication uses dict keys
|
||||
|
|
@ -116,13 +116,12 @@ def initialize_logging(arguments: dict[str, object], asynchronous: bool, route:
|
|||
cb not in litellm._known_custom_logger_compatible_callbacks # pyright: ignore[reportPrivateUsage] # callback compatibility registry has no public accessor
|
||||
or cb in litellm.input_callback + litellm.success_callback + litellm.failure_callback
|
||||
)
|
||||
and cb
|
||||
not in (utils.callback_list or []) # mutable-ok: empty list normalizes an uninitialized callback registry
|
||||
and cb not in (utils.callback_list or ())
|
||||
]
|
||||
if uninitialized:
|
||||
set_callbacks(uninitialized, function_id=arguments.get("id"))
|
||||
utils.callback_list = list( # mutable-ok: global callback registry is mutable
|
||||
dict.fromkeys((*(utils.callback_list or []), *uninitialized))
|
||||
dict.fromkeys((*(utils.callback_list or ()), *uninitialized))
|
||||
) # mutable-ok: global callback registry is mutable
|
||||
if litellm_logging.customLogger is None: # pyright: ignore[reportUnnecessaryComparison] # runtime plugin registry can be reset to None
|
||||
set_callbacks(
|
||||
|
|
@ -183,7 +182,7 @@ def initialize_logging(arguments: dict[str, object], asynchronous: bool, route:
|
|||
supports_correlation_logging=asynchronous,
|
||||
)
|
||||
logger.dynamic_input_callbacks = [ # mutable-ok: remove callbacks promoted to the global registry
|
||||
cb for cb in dict.fromkeys(logger.dynamic_input_callbacks or []) if cb not in litellm.input_callback
|
||||
cb for cb in dict.fromkeys(logger.dynamic_input_callbacks or ()) if cb not in litellm.input_callback
|
||||
]
|
||||
arguments["litellm_call_id"] = call_id
|
||||
arguments["litellm_logging_obj"] = logger
|
||||
|
|
@ -194,7 +193,7 @@ def invoke_terminal(
|
|||
action: str,
|
||||
roots: object,
|
||||
logger: object,
|
||||
record: dict[str, object] | None,
|
||||
record: Mapping[str, object] | None,
|
||||
value: object,
|
||||
fallback_start_time: datetime,
|
||||
fallback_end_time: datetime,
|
||||
|
|
@ -203,13 +202,22 @@ def invoke_terminal(
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
logging: Final = cast(Logging, logger) # cast-ok: Rust passes the logger returned by initialize_logging
|
||||
timing: Final = cast(dict[str, object], record["timing"]) if record is not None else None
|
||||
if not isinstance(logger, Logging):
|
||||
raise TypeError(f"expected Logging, got {type(logger).__name__}")
|
||||
logging: Final = logger
|
||||
timing_value: Final = record.get("timing") if record is not None else None
|
||||
timing: Final = timing_value if isinstance(timing_value, Mapping) else None
|
||||
start_value: Final = timing.get("start_time") if timing is not None else None
|
||||
end_value: Final = timing.get("end_time") if timing is not None else None
|
||||
start_time: Final = (
|
||||
datetime.fromtimestamp(cast(float, timing["start_time"])) if timing is not None else fallback_start_time
|
||||
datetime.fromtimestamp(start_value, tz=fallback_start_time.tzinfo)
|
||||
if isinstance(start_value, (int, float))
|
||||
else fallback_start_time
|
||||
)
|
||||
end_time: Final = (
|
||||
datetime.fromtimestamp(cast(float, timing["end_time"])) if timing is not None else fallback_end_time
|
||||
datetime.fromtimestamp(end_value, tz=fallback_end_time.tzinfo)
|
||||
if isinstance(end_value, (int, float))
|
||||
else fallback_end_time
|
||||
)
|
||||
if action == "sync_success":
|
||||
|
||||
|
|
|
|||
|
|
@ -91,7 +91,7 @@ def attempt(
|
|||
except declined as error:
|
||||
return RustDeclined(reason=_decline_reason(error))
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
raise_upstream(error, context)
|
||||
return RustHandled(adapt(value))
|
||||
|
||||
|
||||
|
|
@ -112,7 +112,7 @@ async def aattempt(
|
|||
except declined as error:
|
||||
return RustDeclined(reason=_decline_reason(error))
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
raise_upstream(error, context)
|
||||
return RustHandled(adapt(value))
|
||||
|
||||
|
||||
|
|
@ -136,7 +136,7 @@ def _required_reason(result: RustDeclined | RustUnavailable) -> str:
|
|||
return f"declined the request: {reason}"
|
||||
|
||||
|
||||
def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn:
|
||||
def raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn:
|
||||
args: Final[tuple[object, ...]] = error.args
|
||||
status_value: Final = args[0] if args else 0
|
||||
message_value: Final = args[1] if len(args) > 1 else str(error)
|
||||
|
|
|
|||
|
|
@ -439,11 +439,14 @@ async def test_fake_stream_context_closes_after_early_cancellation():
|
|||
AnthropicMessagesStreamHiddenParams(additional_headers={}),
|
||||
)
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
async def cancel_inside_context() -> None:
|
||||
async with stream:
|
||||
await anext(stream)
|
||||
raise asyncio.CancelledError
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await cancel_inside_context()
|
||||
|
||||
await stream.aclose()
|
||||
assert completed == [True]
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue