diff --git a/src/utils/compaction.py b/src/utils/compaction.py index 47c2c9530..06d4c1cdb 100644 --- a/src/utils/compaction.py +++ b/src/utils/compaction.py @@ -27,14 +27,16 @@ to disentangle a tangle of side effects. """ +from collections.abc import Callable from datetime import UTC, datetime -from typing import Any +from typing import Any, Optional from ogx_client import AsyncOgxClient from log import get_logger from models.compaction import ConversationSummary from utils.query import normalize_vertex_ai_model_id +from utils.token_counter import TokenCounter from utils.token_estimator import ( estimate_conversation_tokens, estimate_tokens, @@ -44,6 +46,15 @@ logger = get_logger(__name__) +CallCounter = Callable[[str, TokenCounter], None] +"""Receives the model and the token usage of one LLM call made here. + +The provider bills the calls this module makes, so their usage must not be +lost (LCORE-3910). What is done with it (metrics, quota) is the caller's +business: this module only reports each call, as soon as its response +arrived and before the response is looked at. +""" + SUMMARIZATION_PROMPT = ( "Summarize this conversation history for an AI assistant that helps with\n" @@ -197,12 +208,32 @@ def _extract_response_text(response: Any) -> str: return "".join(parts) -async def summarize_chunk( +def reported_usage(response: Any) -> TokenCounter: + """Return the token usage the provider reported for one LLM call. + + Parameters: + response: The result of ``client.responses.create``. + + Returns: + The usage of that one call. The call is counted also when the + provider reported no usage for it; the token counts are 0 then. + """ + usage = getattr(response, "usage", None) + return TokenCounter( + input_tokens=int(getattr(usage, "input_tokens", 0) or 0), + output_tokens=int(getattr(usage, "output_tokens", 0) or 0), + llm_calls=1, + ) + + +async def summarize_chunk( # pylint: disable=too-many-arguments client: AsyncOgxClient, model: str, old_items: list[Any], summarized_through_turn: int, encoding_name: str, + *, + count_call: Optional[CallCounter] = None, ) -> ConversationSummary: """Summarize *old_items* via one LLM call and return a ConversationSummary. @@ -242,6 +273,9 @@ async def summarize_chunk( encoding_name: Tiktoken encoding name used to count tokens in the produced summary. Should match the encoding used to decide the compaction trigger. + count_call: Called with the model and the usage the provider + reported, once the LLM call returned. It is called also + when the call returned no text and this function raises. Returns: A populated ConversationSummary. @@ -275,6 +309,8 @@ async def summarize_chunk( stream=False, store=False, ) + if count_call is not None: + count_call(model, reported_usage(response)) summary_text = _extract_response_text(response).strip() if not summary_text: raise ValueError( @@ -314,6 +350,8 @@ async def recursively_resummarize( model: str, summaries: list[ConversationSummary], encoding_name: str, + *, + count_call: Optional[CallCounter] = None, ) -> ConversationSummary: """Collapse multiple ``ConversationSummary`` records into one. @@ -348,6 +386,9 @@ async def recursively_resummarize( short-circuit before invoking this function. encoding_name: Tiktoken encoding name used to count tokens in the produced fold. + count_call: Called with the model and the usage the provider + reported, once the LLM call returned. It is called also + when the call returned no text and this function raises. Returns: A single ConversationSummary representing the union of @@ -387,6 +428,8 @@ async def recursively_resummarize( stream=False, store=False, ) + if count_call is not None: + count_call(model, reported_usage(response)) folded_text = _extract_response_text(response).strip() if not folded_text: raise ValueError( diff --git a/tests/unit/utils/test_compaction.py b/tests/unit/utils/test_compaction.py index 0eda1dcce..e591f704e 100644 --- a/tests/unit/utils/test_compaction.py +++ b/tests/unit/utils/test_compaction.py @@ -2,7 +2,7 @@ # pylint: disable=too-few-public-methods -from typing import Any +from typing import Any, Optional import pytest from pytest_mock import MockerFixture @@ -17,8 +17,10 @@ is_message_item, partition_conversation, recursively_resummarize, + reported_usage, summarize_chunk, ) +from utils.token_counter import TokenCounter from utils.token_estimator import ( DEFAULT_ENCODING_NAME, estimate_conversation_tokens, @@ -595,3 +597,98 @@ async def test_raises_when_llm_returns_empty(self, mocker: MockerFixture) -> Non summaries=summaries, encoding_name=DEFAULT_ENCODING_NAME, ) + + +# --------------------------------------------------------------------------- +# token usage of the two LLM calls (LCORE-3910) +# --------------------------------------------------------------------------- + +MODEL = "openai/gpt-4o-mini" +CALL_USAGE = TokenCounter(input_tokens=640, output_tokens=72, llm_calls=1) + + +def _make_billed_response(mocker: MockerFixture, text: str) -> Any: + """Build a response that yields *text* and reports ``CALL_USAGE``.""" + response = _make_summary_response(mocker, text) + response.usage = mocker.Mock(input_tokens=640, output_tokens=72) + return response + + +class TestCallUsage: + """What summarize_chunk and recursively_resummarize report of their call.""" + + def test_reported_usage(self, mocker: MockerFixture) -> None: + """The usage of a call is what the provider reported for it.""" + assert reported_usage(_make_billed_response(mocker, "text")) == CALL_USAGE + + @pytest.mark.parametrize( + "usage", + [None, {"input_tokens": None, "output_tokens": None}, {}], + ids=["no usage", "empty counts", "no counts"], + ) + def test_call_without_reported_usage_is_still_a_call( + self, mocker: MockerFixture, usage: Optional[dict[str, Any]] + ) -> None: + """Without reported usage the counts stay at zero; the call counts.""" + response = mocker.Mock( + usage=None if usage is None else mocker.Mock(spec_set=list(usage), **usage) + ) + assert reported_usage(response) == TokenCounter(llm_calls=1) + + @pytest.mark.asyncio + async def test_summarize_chunk_reports_its_call( + self, mocker: MockerFixture + ) -> None: + """The summarization call is reported with the usage the provider gave.""" + client = mocker.AsyncMock() + client.responses.create.return_value = _make_billed_response(mocker, "text") + count_call = mocker.Mock() + await summarize_chunk( + client=client, + model=MODEL, + old_items=[_MessageItem("user", "hi")], + summarized_through_turn=1, + encoding_name=DEFAULT_ENCODING_NAME, + count_call=count_call, + ) + count_call.assert_called_once_with(MODEL, CALL_USAGE) + + @pytest.mark.asyncio + async def test_summarize_chunk_reports_a_call_that_returned_no_text( + self, mocker: MockerFixture + ) -> None: + """The call was made and billed, also when its answer holds no summary.""" + client = mocker.AsyncMock() + client.responses.create.return_value = _make_billed_response(mocker, "") + count_call = mocker.Mock() + with pytest.raises(ValueError, match="no extractable text"): + await summarize_chunk( + client=client, + model=MODEL, + old_items=[_MessageItem("user", "hi")], + summarized_through_turn=1, + encoding_name=DEFAULT_ENCODING_NAME, + count_call=count_call, + ) + count_call.assert_called_once_with(MODEL, CALL_USAGE) + + @pytest.mark.asyncio + async def test_fold_reports_a_call_that_returned_no_text( + self, mocker: MockerFixture + ) -> None: + """The fold call was made and billed, also when its answer is empty.""" + client = mocker.AsyncMock() + client.responses.create.return_value = _make_billed_response(mocker, "") + count_call = mocker.Mock() + with pytest.raises(ValueError, match="no extractable text"): + await recursively_resummarize( + client=client, + model=MODEL, + summaries=[ + _make_summary("a", through=1), + _make_summary("b", through=2), + ], + encoding_name=DEFAULT_ENCODING_NAME, + count_call=count_call, + ) + count_call.assert_called_once_with(MODEL, CALL_USAGE)