mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
Merge pull request #5480 from BerriAI/litellm_track_streaming_spendLogs
[Feat] Track Usage for `/streamGenerateContent` endpoint
This commit is contained in:
commit
a64f9f4bc0
5 changed files with 193 additions and 7 deletions
|
|
@ -22,6 +22,9 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.vertex_ai_and_google_ai_studio.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
ModelResponseIterator,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
ConfigFieldInfo,
|
||||
ConfigFieldUpdate,
|
||||
|
|
@ -32,7 +35,9 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
from .streaming_handler import chunk_processor
|
||||
from .success_handler import PassThroughEndpointLogging
|
||||
from .types import EndpointType
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
|
@ -284,6 +289,12 @@ def get_response_headers(headers: httpx.Headers) -> dict:
|
|||
return return_headers
|
||||
|
||||
|
||||
def get_endpoint_type(url: str) -> EndpointType:
|
||||
if ("generateContent") in url or ("streamGenerateContent") in url:
|
||||
return EndpointType.VERTEX_AI
|
||||
return EndpointType.GENERIC
|
||||
|
||||
|
||||
async def pass_through_request(
|
||||
request: Request,
|
||||
target: str,
|
||||
|
|
@ -307,6 +318,8 @@ async def pass_through_request(
|
|||
request=request, headers=headers, forward_headers=forward_headers
|
||||
)
|
||||
|
||||
endpoint_type: EndpointType = get_endpoint_type(str(url))
|
||||
|
||||
_parsed_body = None
|
||||
if custom_body:
|
||||
_parsed_body = custom_body
|
||||
|
|
@ -416,9 +429,15 @@ async def pass_through_request(
|
|||
status_code=e.response.status_code, detail=await e.response.aread()
|
||||
)
|
||||
|
||||
# Create an async generator to yield the response content
|
||||
async def stream_response() -> AsyncIterable[bytes]:
|
||||
async for chunk in response.aiter_bytes():
|
||||
async for chunk in chunk_processor(
|
||||
response.aiter_bytes(),
|
||||
litellm_logging_obj=logging_obj,
|
||||
endpoint_type=endpoint_type,
|
||||
start_time=start_time,
|
||||
passthrough_success_handler_obj=pass_through_endpoint_logging,
|
||||
url_route=str(url),
|
||||
):
|
||||
yield chunk
|
||||
|
||||
return StreamingResponse(
|
||||
|
|
@ -454,10 +473,15 @@ async def pass_through_request(
|
|||
status_code=e.response.status_code, detail=await e.response.aread()
|
||||
)
|
||||
|
||||
# streaming response
|
||||
# Create an async generator to yield the response content
|
||||
async def stream_response() -> AsyncIterable[bytes]:
|
||||
async for chunk in response.aiter_bytes():
|
||||
async for chunk in chunk_processor(
|
||||
response.aiter_bytes(),
|
||||
litellm_logging_obj=logging_obj,
|
||||
endpoint_type=endpoint_type,
|
||||
start_time=start_time,
|
||||
passthrough_success_handler_obj=pass_through_endpoint_logging,
|
||||
url_route=str(url),
|
||||
):
|
||||
yield chunk
|
||||
|
||||
return StreamingResponse(
|
||||
|
|
|
|||
117
litellm/proxy/pass_through_endpoints/streaming_handler.py
Normal file
117
litellm/proxy/pass_through_endpoints/streaming_handler.py
Normal file
|
|
@ -0,0 +1,117 @@
|
|||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from typing import AsyncIterable, Dict, List, Optional, Union
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.vertex_ai_and_google_ai_studio.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
ModelResponseIterator as VertexAIIterator,
|
||||
)
|
||||
from litellm.types.utils import GenericStreamingChunk
|
||||
|
||||
from .success_handler import PassThroughEndpointLogging
|
||||
from .types import EndpointType
|
||||
|
||||
|
||||
def get_litellm_chunk(
|
||||
model_iterator: VertexAIIterator,
|
||||
custom_stream_wrapper: litellm.utils.CustomStreamWrapper,
|
||||
chunk_dict: Dict,
|
||||
) -> Optional[Dict]:
|
||||
|
||||
generic_chunk: GenericStreamingChunk = model_iterator.chunk_parser(chunk_dict)
|
||||
if generic_chunk:
|
||||
return custom_stream_wrapper.chunk_creator(chunk=generic_chunk)
|
||||
return None
|
||||
|
||||
|
||||
def get_iterator_class_from_endpoint_type(
|
||||
endpoint_type: EndpointType,
|
||||
) -> Optional[type]:
|
||||
if endpoint_type == EndpointType.VERTEX_AI:
|
||||
return VertexAIIterator
|
||||
return None
|
||||
|
||||
|
||||
async def chunk_processor(
|
||||
aiter_bytes: AsyncIterable[bytes],
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
endpoint_type: EndpointType,
|
||||
start_time: datetime,
|
||||
passthrough_success_handler_obj: PassThroughEndpointLogging,
|
||||
url_route: str,
|
||||
) -> AsyncIterable[bytes]:
|
||||
|
||||
iteratorClass = get_iterator_class_from_endpoint_type(endpoint_type)
|
||||
if iteratorClass is None:
|
||||
# Generic endpoint - litellm does not do any tracking / logging for this
|
||||
async for chunk in aiter_bytes:
|
||||
yield chunk
|
||||
else:
|
||||
# known streaming endpoint - litellm will do tracking / logging for this
|
||||
model_iterator = iteratorClass(
|
||||
sync_stream=False, streaming_response=aiter_bytes
|
||||
)
|
||||
custom_stream_wrapper = litellm.utils.CustomStreamWrapper(
|
||||
completion_stream=aiter_bytes, model=None, logging_obj=litellm_logging_obj
|
||||
)
|
||||
buffer = b""
|
||||
all_chunks = []
|
||||
async for chunk in aiter_bytes:
|
||||
buffer += chunk
|
||||
try:
|
||||
_decoded_chunk = chunk.decode("utf-8")
|
||||
_chunk_dict = json.loads(_decoded_chunk)
|
||||
litellm_chunk = get_litellm_chunk(
|
||||
model_iterator, custom_stream_wrapper, _chunk_dict
|
||||
)
|
||||
if litellm_chunk:
|
||||
all_chunks.append(litellm_chunk)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
finally:
|
||||
yield chunk # Yield the original bytes
|
||||
|
||||
# Process any remaining data in the buffer
|
||||
if buffer:
|
||||
try:
|
||||
_chunk_dict = json.loads(buffer.decode("utf-8"))
|
||||
|
||||
if isinstance(_chunk_dict, list):
|
||||
for _chunk in _chunk_dict:
|
||||
litellm_chunk = get_litellm_chunk(
|
||||
model_iterator, custom_stream_wrapper, _chunk
|
||||
)
|
||||
if litellm_chunk:
|
||||
all_chunks.append(litellm_chunk)
|
||||
elif isinstance(_chunk_dict, dict):
|
||||
litellm_chunk = get_litellm_chunk(
|
||||
model_iterator, custom_stream_wrapper, _chunk_dict
|
||||
)
|
||||
if litellm_chunk:
|
||||
all_chunks.append(litellm_chunk)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
complete_streaming_response: Optional[
|
||||
Union[litellm.ModelResponse, litellm.TextCompletionResponse]
|
||||
] = litellm.stream_chunk_builder(chunks=all_chunks)
|
||||
if complete_streaming_response is None:
|
||||
complete_streaming_response = litellm.ModelResponse()
|
||||
end_time = datetime.now()
|
||||
|
||||
if passthrough_success_handler_obj.is_vertex_route(url_route):
|
||||
_model = passthrough_success_handler_obj.extract_model_from_url(url_route)
|
||||
complete_streaming_response.model = _model
|
||||
litellm_logging_obj.model = _model
|
||||
litellm_logging_obj.model_call_details["model"] = _model
|
||||
|
||||
asyncio.create_task(
|
||||
litellm_logging_obj.async_success_handler(
|
||||
result=complete_streaming_response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
)
|
||||
6
litellm/proxy/pass_through_endpoints/types.py
Normal file
6
litellm/proxy/pass_through_endpoints/types.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from enum import Enum
|
||||
|
||||
|
||||
class EndpointType(str, Enum):
|
||||
VERTEX_AI = "vertex-ai"
|
||||
GENERIC = "generic"
|
||||
|
|
@ -10,7 +10,12 @@ vertexai.init(
|
|||
api_transport="rest",
|
||||
)
|
||||
|
||||
model = GenerativeModel(model_name="gemini-1.0-pro")
|
||||
response = model.generate_content("hi")
|
||||
model = GenerativeModel(model_name="gemini-1.5-flash-001")
|
||||
response = model.generate_content(
|
||||
"hi tell me a joke and a very long story", stream=True
|
||||
)
|
||||
|
||||
print("response", response)
|
||||
|
||||
for chunk in response:
|
||||
print(chunk)
|
||||
|
|
|
|||
|
|
@ -117,3 +117,37 @@ async def test_basic_vertex_ai_pass_through_with_spendlog():
|
|||
)
|
||||
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_basic_vertex_ai_pass_through_streaming_with_spendlog():
|
||||
|
||||
spend_before = await call_spend_logs_endpoint() or 0.0
|
||||
print("spend_before", spend_before)
|
||||
load_vertex_ai_credentials()
|
||||
|
||||
vertexai.init(
|
||||
project="adroit-crow-413218",
|
||||
location="us-central1",
|
||||
api_endpoint=f"{LITE_LLM_ENDPOINT}/vertex-ai",
|
||||
api_transport="rest",
|
||||
)
|
||||
|
||||
model = GenerativeModel(model_name="gemini-1.0-pro")
|
||||
response = model.generate_content("hi", stream=True)
|
||||
|
||||
for chunk in response:
|
||||
print("chunk", chunk)
|
||||
|
||||
print("response", response)
|
||||
|
||||
await asyncio.sleep(20)
|
||||
spend_after = await call_spend_logs_endpoint()
|
||||
print("spend_after", spend_after)
|
||||
assert (
|
||||
spend_after > spend_before
|
||||
), "Spend should be greater than before. spend_before: {}, spend_after: {}".format(
|
||||
spend_before, spend_after
|
||||
)
|
||||
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue