mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
perf(proxy): keep stream end output and Redis cache debug strings off the event loop (#45252)
* perf(proxy): keep end-of-stream output assembly and Redis pipeline logging off the event loop Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * perf(caching): serialize Redis pipeline values with orjson Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * perf(proxy): resolve model-level guardrails once per streamed request Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): type the streamed guardrail cache Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): fall back to json.dumps when orjson is not installed Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(caching): cover the json.dumps fallback for values orjson rejects Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep new guardrail cache and orjson helper within the type gates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * perf(caching): guard Redis print_verbose value interpolation behind is_debugging_on Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * revert(caching): drop the orjson fast path from Redis pipeline writes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(caching): drop the debug-on print_verbose test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * revert(proxy): drop the per-request guardrail lookup cache from this PR Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * perf(logging): let print_verbose take %s args so Redis cache values format only when verbose Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): type set_cache key/value and redis_version so print_verbose args are known to pyright Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry <kerry@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
231627ecd9
commit
dcd0630e98
6 changed files with 160 additions and 21 deletions
|
|
@ -1096,10 +1096,12 @@ def _enable_debugging():
|
|||
verbose_proxy_stdout_logger.disabled = False
|
||||
|
||||
|
||||
def print_verbose(print_statement):
|
||||
def print_verbose(print_statement: object, *args: object) -> None:
|
||||
try:
|
||||
if set_verbose:
|
||||
print(redact_secrets(str(print_statement))) # noqa: T201
|
||||
if not set_verbose:
|
||||
return
|
||||
message: Final = str(print_statement) % args if args else str(print_statement)
|
||||
print(redact_secrets(message)) # noqa: T201
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
|
|
|||
|
|
@ -705,7 +705,7 @@ class RedisCache(BaseCache):
|
|||
self.redis_flush_size: int = 100
|
||||
else:
|
||||
self.redis_flush_size = redis_flush_size
|
||||
self.redis_version = "Unknown"
|
||||
self.redis_version: str = "Unknown"
|
||||
try:
|
||||
if not coroutine_checker.is_async_callable(self.redis_client):
|
||||
self.redis_version = self.redis_client.info()["redis_version"]
|
||||
|
|
@ -879,9 +879,15 @@ class RedisCache(BaseCache):
|
|||
# Fallback for unparseable versions (e.g., "v7.0.0", "latest")
|
||||
return DEFAULT_REDIS_MAJOR_VERSION
|
||||
|
||||
def set_cache(self, key, value, **kwargs):
|
||||
def set_cache(self, key: str, value: object, **kwargs):
|
||||
ttl: Final = self.get_ttl(**kwargs)
|
||||
print_verbose(f"Set Redis Cache: key: {key}\nValue {value}\nttl={ttl}, redis_version={self.redis_version}")
|
||||
print_verbose(
|
||||
"Set Redis Cache: key: %s\nValue %s\nttl=%s, redis_version=%s",
|
||||
key,
|
||||
value,
|
||||
ttl,
|
||||
self.redis_version,
|
||||
)
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
try:
|
||||
start_time: Final = time.time()
|
||||
|
|
@ -1134,7 +1140,7 @@ class RedisCache(BaseCache):
|
|||
raise ValueError("Redis client does not support Lua script registration")
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_set_cache(self, key, value, **kwargs):
|
||||
async def async_set_cache(self, key: str | None, value: object, **kwargs):
|
||||
from redis.asyncio import Redis
|
||||
|
||||
if key is None:
|
||||
|
|
@ -1170,7 +1176,7 @@ class RedisCache(BaseCache):
|
|||
key = self.check_and_fix_namespace(key=key)
|
||||
ttl: Final = self.get_ttl(**kwargs)
|
||||
nx: Final = kwargs.get("nx", False)
|
||||
print_verbose(f"Set ASYNC Redis Cache: key: {key}\nValue {value}\nttl={ttl}")
|
||||
print_verbose("Set ASYNC Redis Cache: key: %s\nValue %s\nttl=%s", key, value, ttl)
|
||||
|
||||
try:
|
||||
if not hasattr(_redis_client, "set"):
|
||||
|
|
@ -1181,7 +1187,7 @@ class RedisCache(BaseCache):
|
|||
nx=nx,
|
||||
ex=ttl,
|
||||
)
|
||||
print_verbose(f"Successfully Set ASYNC Redis Cache: key: {key}\nValue {value}\nttl={ttl}")
|
||||
print_verbose("Successfully Set ASYNC Redis Cache: key: %s\nValue %s\nttl=%s", key, value, ttl)
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
asyncio.create_task(
|
||||
|
|
@ -1229,7 +1235,7 @@ class RedisCache(BaseCache):
|
|||
# Iterate through each key-value pair in the cache_list and set them in the pipeline.
|
||||
for cache_key, cache_value in cache_list:
|
||||
cache_key = self.check_and_fix_namespace(key=cache_key)
|
||||
print_verbose(f"Set ASYNC Redis Cache PIPELINE: key: {cache_key}\nValue {cache_value}\nttl={ttl}")
|
||||
print_verbose("Set ASYNC Redis Cache PIPELINE: key: %s\nValue %s\nttl=%s", cache_key, cache_value, ttl)
|
||||
json_cache_value = json.dumps(cache_value)
|
||||
# Set the value with a TTL if it's provided.
|
||||
_td: timedelta | None = None
|
||||
|
|
@ -1258,7 +1264,12 @@ class RedisCache(BaseCache):
|
|||
_redis_client: Final = self.init_async_client()
|
||||
start_time: Final = time.time()
|
||||
|
||||
print_verbose(f"Set Async Redis Cache: key list: {cache_list}\nttl={ttl}, redis_version={self.redis_version}")
|
||||
print_verbose(
|
||||
"Set Async Redis Cache: key list: %s\nttl=%s, redis_version=%s",
|
||||
cache_list,
|
||||
ttl,
|
||||
self.redis_version,
|
||||
)
|
||||
try:
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
results: Final = await self._pipeline_helper(pipe, cache_list, ttl)
|
||||
|
|
@ -1396,10 +1407,10 @@ class RedisCache(BaseCache):
|
|||
raise e
|
||||
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
print_verbose(f"Set ASYNC Redis Cache: key: {key}\nValue {value}\nttl={ttl}")
|
||||
print_verbose("Set ASYNC Redis Cache: key: %s\nValue %s\nttl=%s", key, value, ttl)
|
||||
try:
|
||||
await self._set_cache_sadd_helper(redis_client=_redis_client, key=key, value=value, ttl=ttl)
|
||||
print_verbose(f"Successfully Set ASYNC Redis Cache SADD: key: {key}\nValue {value}\nttl={ttl}")
|
||||
print_verbose("Successfully Set ASYNC Redis Cache SADD: key: %s\nValue %s\nttl=%s", key, value, ttl)
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
asyncio.create_task(
|
||||
|
|
@ -1605,7 +1616,7 @@ class RedisCache(BaseCache):
|
|||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
print_verbose(f"Got Redis Cache: key: {key}, cached_response {cached_response}")
|
||||
print_verbose("Got Redis Cache: key: %s, cached_response %s", key, cached_response)
|
||||
return self._get_cache_logic(cached_response=cached_response)
|
||||
except Exception as e:
|
||||
log_redis_failure(
|
||||
|
|
@ -1705,7 +1716,7 @@ class RedisCache(BaseCache):
|
|||
try:
|
||||
print_verbose(f"Get Async Redis Cache: key: {key}")
|
||||
cached_response: Final = await _redis_client.get(key)
|
||||
print_verbose(f"Got Async Redis Cache: key: {key}, cached_response {cached_response}")
|
||||
print_verbose("Got Async Redis Cache: key: %s, cached_response %s", key, cached_response)
|
||||
response: Final = self._get_cache_logic(cached_response=cached_response)
|
||||
|
||||
end_time = time.time()
|
||||
|
|
@ -2019,7 +2030,7 @@ class RedisCache(BaseCache):
|
|||
_redis_client: Final[Redis] = self.init_async_client()
|
||||
start_time: Final = time.time()
|
||||
|
||||
print_verbose(f"Increment Async Redis Cache Pipeline: increment list: {increment_list}")
|
||||
print_verbose("Increment Async Redis Cache Pipeline: increment list: %s", increment_list)
|
||||
|
||||
try:
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
|
|
|
|||
|
|
@ -4075,10 +4075,10 @@ class ProxyLogging:
|
|||
yield chunk
|
||||
except (GeneratorExit, asyncio.CancelledError):
|
||||
await ProxyLogging._close_guarded_layers(guarded_layers)
|
||||
ProxyLogging._record_served_stream_output(request_data, served_chunks)
|
||||
await ProxyLogging._record_served_stream_output(request_data, served_chunks)
|
||||
raise
|
||||
except Exception as e:
|
||||
ProxyLogging._record_served_stream_output(request_data, served_chunks)
|
||||
await ProxyLogging._record_served_stream_output(request_data, served_chunks)
|
||||
if not ProxyLogging._discard_deferred_stream_logging_for_failure(request_data, e):
|
||||
ProxyLogging.fire_deferred_stream_logging(request_data)
|
||||
raise
|
||||
|
|
@ -4087,7 +4087,7 @@ class ProxyLogging:
|
|||
# completed. unified_guardrail writes guardrail_information during
|
||||
# its end-of-stream block (inside current_response), so by the time
|
||||
# we reach this point the metadata is fully populated.
|
||||
ProxyLogging._record_served_stream_output(request_data, served_chunks)
|
||||
await ProxyLogging._record_served_stream_output(request_data, served_chunks)
|
||||
ProxyLogging.fire_deferred_stream_logging(request_data)
|
||||
|
||||
async def _pipeline_gated_stream(
|
||||
|
|
@ -4166,11 +4166,12 @@ class ProxyLogging:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _record_served_stream_output(request_data: Mapping[str, object], served_chunks: Sequence[object]) -> None:
|
||||
async def _record_served_stream_output(request_data: Mapping[str, object], served_chunks: Sequence[object]) -> None:
|
||||
logging_obj: Final = request_data.get("litellm_logging_obj")
|
||||
if not isinstance(logging_obj, Logging):
|
||||
return
|
||||
record_served_output_texts(logging_obj.model_call_details, served_stream_output_texts(served_chunks))
|
||||
texts: Final = await offload_token_count(served_stream_output_texts)(served_chunks)
|
||||
record_served_output_texts(logging_obj.model_call_details, texts)
|
||||
|
||||
@staticmethod
|
||||
def fire_deferred_stream_logging(request_data: dict) -> None:
|
||||
|
|
|
|||
|
|
@ -1611,3 +1611,27 @@ def test_call_stack_info_skips_native_lifecycle_frames():
|
|||
return native_drive()
|
||||
|
||||
assert anthropic_messages() == "anthropic_messages <- test_call_stack_info_skips_native_lifecycle_frames"
|
||||
|
||||
|
||||
class _FormatCountingStr(str):
|
||||
format_calls = 0
|
||||
|
||||
def __format__(self, spec: str) -> str:
|
||||
_FormatCountingStr.format_calls += 1
|
||||
return super().__format__(spec)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_set_cache_pipeline_does_not_format_cached_values_for_logging(monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
pipe = _SetRecordingPipeline()
|
||||
client = MagicMock()
|
||||
client.pipeline = MagicMock(return_value=pipe)
|
||||
|
||||
_FormatCountingStr.format_calls = 0
|
||||
with patch.object(redis_cache, "init_async_client", return_value=client):
|
||||
await redis_cache.async_set_cache_pipeline([("k1", _FormatCountingStr("embedding"))], ttl=60)
|
||||
|
||||
assert _FormatCountingStr.format_calls == 0
|
||||
assert pipe.sets == [("k1", '"embedding"', timedelta(seconds=60))]
|
||||
|
|
|
|||
|
|
@ -3327,6 +3327,83 @@ def test_handle_exception_on_proxy_preserves_auth_error_status_code():
|
|||
assert int(result.code) == 401, f"Expected 401, got {result.code}"
|
||||
|
||||
|
||||
def _anthropic_thinking_sse(deltas: int) -> list[bytes]:
|
||||
def event(name: str, payload: dict) -> bytes:
|
||||
return f"event: {name}\ndata: {json.dumps(payload)}\n\n".encode()
|
||||
|
||||
message: Final = {"id": "msg_1", "type": "message", "role": "assistant", "model": "claude-sonnet-4-5"}
|
||||
return [
|
||||
event(
|
||||
"message_start",
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {**message, "content": [], "usage": {"input_tokens": 1, "output_tokens": 1}},
|
||||
},
|
||||
),
|
||||
event(
|
||||
"content_block_start",
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "thinking", "thinking": ""}},
|
||||
),
|
||||
*(
|
||||
event(
|
||||
"content_block_delta",
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "thinking_delta", "thinking": f"t{i} "}},
|
||||
)
|
||||
for i in range(deltas)
|
||||
),
|
||||
event("content_block_stop", {"type": "content_block_stop", "index": 0}),
|
||||
event(
|
||||
"content_block_start",
|
||||
{"type": "content_block_start", "index": 1, "content_block": {"type": "text", "text": ""}},
|
||||
),
|
||||
event(
|
||||
"content_block_delta",
|
||||
{"type": "content_block_delta", "index": 1, "delta": {"type": "text_delta", "text": "final answer"}},
|
||||
),
|
||||
event("content_block_stop", {"type": "content_block_stop", "index": 1}),
|
||||
event(
|
||||
"message_delta",
|
||||
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 1}},
|
||||
),
|
||||
event("message_stop", {"type": "message_stop"}),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_record_served_stream_output_lets_the_event_loop_run_while_assembling_the_stream():
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.served_output_texts import SERVED_OUTPUT_TEXTS_KEY
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
logging_obj: Final = Logging(
|
||||
model="claude-sonnet-4-5",
|
||||
messages=[],
|
||||
stream=True,
|
||||
call_type="anthropic_messages",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="served-output",
|
||||
function_id="served-output",
|
||||
)
|
||||
loop_turns: Final = []
|
||||
done: Final = asyncio.Event()
|
||||
|
||||
async def count_loop_turns() -> None:
|
||||
while not done.is_set():
|
||||
loop_turns.append(None)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
counter: Final = asyncio.create_task(count_loop_turns())
|
||||
await asyncio.sleep(0)
|
||||
turns_before: Final = len(loop_turns)
|
||||
await ProxyLogging._record_served_stream_output({"litellm_logging_obj": logging_obj}, _anthropic_thinking_sse(2000))
|
||||
turns_during: Final = len(loop_turns) - turns_before
|
||||
done.set()
|
||||
await counter
|
||||
|
||||
assert logging_obj.model_call_details[SERVED_OUTPUT_TEXTS_KEY] == ("final answer",)
|
||||
assert turns_during > 0
|
||||
|
||||
|
||||
class _RedisDown:
|
||||
async def async_delete_cache(self, key: str) -> None:
|
||||
raise ConnectionError("redis is down")
|
||||
|
|
|
|||
|
|
@ -1766,3 +1766,27 @@ def test_diagnostic_filter_redacts_a_non_string_message_object(monkeypatch, nati
|
|||
assert DiagnosticProcessingFilter().filter(record) is True
|
||||
|
||||
assert secret not in record.getMessage()
|
||||
|
||||
|
||||
class _StrCountingValue:
|
||||
str_calls = 0
|
||||
|
||||
def __str__(self) -> str:
|
||||
_StrCountingValue.str_calls += 1
|
||||
return "cached-vector"
|
||||
|
||||
|
||||
def test_print_verbose_formats_args_only_when_set_verbose_is_on(monkeypatch, capsys):
|
||||
import litellm._logging as logging_module
|
||||
|
||||
_StrCountingValue.str_calls = 0
|
||||
|
||||
monkeypatch.setattr(logging_module, "set_verbose", False)
|
||||
logging_module.print_verbose("key %s value %s", "k1", _StrCountingValue())
|
||||
assert _StrCountingValue.str_calls == 0
|
||||
assert capsys.readouterr().out == ""
|
||||
|
||||
monkeypatch.setattr(logging_module, "set_verbose", True)
|
||||
logging_module.print_verbose("key %s value %s", "k1", _StrCountingValue())
|
||||
assert _StrCountingValue.str_calls == 1
|
||||
assert capsys.readouterr().out == "key k1 value cached-vector\n"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue