Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 45 additions & 2 deletions src/utils/compaction.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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"
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
99 changes: 98 additions & 1 deletion tests/unit/utils/test_compaction.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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)
Loading