diff --git a/holmes/core/llm.py b/holmes/core/llm.py index 0ea0461500..2c8e8451a1 100644 --- a/holmes/core/llm.py +++ b/holmes/core/llm.py @@ -3,7 +3,7 @@ from abc import abstractmethod from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING -from litellm.types.utils import ModelResponse +from litellm.types.utils import ModelResponse, TextCompletionResponse import sentry_sdk from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper @@ -530,3 +530,25 @@ def _create_model_entry( "is_robusta_model": is_robusta_model, "model": model, } + + +def get_llm_usage( + llm_response: Union[ModelResponse, CustomStreamWrapper, TextCompletionResponse], +) -> dict: + usage: dict = {} + if ( + ( + isinstance(llm_response, ModelResponse) + or isinstance(llm_response, TextCompletionResponse) + ) + and hasattr(llm_response, "usage") + and llm_response.usage + ): # type: ignore + usage["prompt_tokens"] = llm_response.usage.prompt_tokens # type: ignore + usage["completion_tokens"] = llm_response.usage.completion_tokens # type: ignore + usage["total_tokens"] = llm_response.usage.total_tokens # type: ignore + elif isinstance(llm_response, CustomStreamWrapper): + complete_response = litellm.stream_chunk_builder(chunks=llm_response) # type: ignore + if complete_response: + return get_llm_usage(complete_response) + return usage diff --git a/holmes/core/tool_calling_llm.py b/holmes/core/tool_calling_llm.py index 4077436046..2c630c11c0 100644 --- a/holmes/core/tool_calling_llm.py +++ b/holmes/core/tool_calling_llm.py @@ -27,7 +27,7 @@ is_response_an_incorrect_tool_call, ) from holmes.core.issue import Issue -from holmes.core.llm import LLM +from holmes.core.llm import LLM, get_llm_usage from holmes.core.performance_timing import PerformanceTiming from holmes.core.resource_instruction import ResourceInstructions from holmes.core.runbooks import RunbookManager @@ -422,7 +422,9 @@ def call( # type: ignore ) costs.total_cost += post_processing_cost + self.llm.count_tokens_for_message(messages) perf_timing.end(f"- completed in {i} iterations -") + metadata["usage"] = get_llm_usage(full_response) return LLMResult( result=post_processed_response, unprocessed_result=raw_response, @@ -863,6 +865,8 @@ def call_stream( tools_to_call = getattr(response_message, "tool_calls", None) if not tools_to_call: + self.llm.count_tokens_for_message(messages) + metadata["usage"] = get_llm_usage(full_response) yield StreamMessage( event=StreamEvents.ANSWER_END, data={