Skip to content
Closed
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
85 changes: 56 additions & 29 deletions instructor/v2/core/usage.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,10 @@
from __future__ import annotations

import logging
from typing import TYPE_CHECKING, TypeVar
from numbers import Real
from typing import TYPE_CHECKING, TypeVar, cast

from pydantic import BaseModel

if TYPE_CHECKING:
from anthropic.types import Usage as AnthropicUsage
Expand All @@ -13,6 +16,57 @@
T_Response = TypeVar("T_Response")


def _field_names(model: BaseModel) -> set[str]:
"""Return declared and Pydantic extra field names for a model."""
return set(type(model).model_fields) | set(model.model_extra or {})


def _all_field_names(*models: BaseModel) -> set[str]:
return {field_name for model in models for field_name in _field_names(model)}


def _is_numeric(value: object) -> bool:
return isinstance(value, Real) and not isinstance(value, bool)


def _zero_numeric_fields(model: BaseModel) -> BaseModel:
for field_name in _field_names(model):
value = getattr(model, field_name, None)
if isinstance(value, BaseModel):
_zero_numeric_fields(value)
elif _is_numeric(value):
setattr(model, field_name, 0)
return model


def _accumulate_models(response: BaseModel, total: BaseModel) -> None:
for field_name in _all_field_names(response, total):
response_value = getattr(response, field_name, None)
total_value = getattr(total, field_name, None)
if isinstance(response_value, BaseModel):
if total_value is None:
total_value = _zero_numeric_fields(response_value.model_copy(deep=True))
setattr(total, field_name, total_value)
if isinstance(total_value, BaseModel):
_accumulate_models(response_value, total_value)
setattr(response, field_name, total_value.model_copy(deep=True))
elif isinstance(total_value, BaseModel):
setattr(response, field_name, total_value.model_copy(deep=True))
elif _is_numeric(response_value):
response_number = cast(Real, response_value)
total_number = cast(Real, total_value) if _is_numeric(total_value) else 0
value = total_number + response_number
setattr(total, field_name, value)
setattr(response, field_name, value)
elif _is_numeric(total_value):
setattr(response, field_name, total_value)
elif response_value is not None:
setattr(total, field_name, response_value)
setattr(response, field_name, response_value)
elif total_value is not None:
setattr(response, field_name, total_value)


def update_total_usage(
response: T_Response | None,
total_usage: OpenAIUsage | AnthropicUsage,
Expand All @@ -26,34 +80,7 @@ def update_total_usage(
if isinstance(response_usage, _OpenAIUsage) and isinstance(
total_usage, _OpenAIUsage
):
total_usage.completion_tokens += response_usage.completion_tokens or 0
total_usage.prompt_tokens += response_usage.prompt_tokens or 0
total_usage.total_tokens += response_usage.total_tokens or 0
if (rtd := response_usage.completion_tokens_details) and (
ttd := total_usage.completion_tokens_details
):
ttd.audio_tokens = (ttd.audio_tokens or 0) + (rtd.audio_tokens or 0)
ttd.reasoning_tokens = (ttd.reasoning_tokens or 0) + (
rtd.reasoning_tokens or 0
)
if (rpd := response_usage.prompt_tokens_details) and (
tpd := total_usage.prompt_tokens_details
):
tpd.audio_tokens = (tpd.audio_tokens or 0) + (rpd.audio_tokens or 0)
tpd.cached_tokens = (tpd.cached_tokens or 0) + (rpd.cached_tokens or 0)
response_usage.completion_tokens = total_usage.completion_tokens
response_usage.prompt_tokens = total_usage.prompt_tokens
response_usage.total_tokens = total_usage.total_tokens
response_usage.completion_tokens_details = (
total_usage.completion_tokens_details.model_copy(deep=True)
if total_usage.completion_tokens_details is not None
else None
)
response_usage.prompt_tokens_details = (
total_usage.prompt_tokens_details.model_copy(deep=True)
if total_usage.prompt_tokens_details is not None
else None
)
_accumulate_models(response_usage, total_usage)
return response

try:
Expand Down
17 changes: 3 additions & 14 deletions instructor/v2/providers/anthropic/usage.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@

from typing import Any

from instructor.v2.core.usage import _accumulate_models


def initialize_usage() -> Any:
"""Create an empty Anthropic usage accumulator."""
Expand All @@ -26,18 +28,5 @@ def update_total_usage(response_usage: Any, total_usage: Any) -> bool:
):
return False

if not total_usage.cache_creation_input_tokens:
total_usage.cache_creation_input_tokens = 0
if not total_usage.cache_read_input_tokens:
total_usage.cache_read_input_tokens = 0
total_usage.input_tokens += response_usage.input_tokens or 0
total_usage.output_tokens += response_usage.output_tokens or 0
total_usage.cache_creation_input_tokens += (
response_usage.cache_creation_input_tokens or 0
)
total_usage.cache_read_input_tokens += response_usage.cache_read_input_tokens or 0
response_usage.input_tokens = total_usage.input_tokens
response_usage.output_tokens = total_usage.output_tokens
response_usage.cache_creation_input_tokens = total_usage.cache_creation_input_tokens
response_usage.cache_read_input_tokens = total_usage.cache_read_input_tokens
_accumulate_models(response_usage, total_usage)
return True
33 changes: 33 additions & 0 deletions tests/coverage/test_anthropic_support_coverage.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,8 @@
import httpx
import pytest
from anthropic.types import Usage
from anthropic.types.cache_creation import CacheCreation
from anthropic.types.server_tool_usage import ServerToolUsage
from pydantic import BaseModel, ValidationInfo, field_validator

from instructor.v2.core.client import AsyncInstructor, Instructor
Expand Down Expand Up @@ -397,3 +399,34 @@ def test_anthropic_usage_initializes_and_accumulates_sdk_usage() -> None:
) == (12, 10, 2, 5)
assert update_total_usage(object(), total) is False
assert update_total_usage(second, object()) is False


def test_anthropic_usage_accumulates_nested_provider_fields() -> None:
total = initialize_usage()
first = Usage(
input_tokens=100,
output_tokens=50,
cache_creation=CacheCreation(
ephemeral_1h_input_tokens=10, ephemeral_5m_input_tokens=20
),
server_tool_use=ServerToolUsage(web_fetch_requests=2, web_search_requests=3),
)
second = Usage(
input_tokens=100,
output_tokens=50,
cache_creation=CacheCreation(
ephemeral_1h_input_tokens=10, ephemeral_5m_input_tokens=20
),
server_tool_use=ServerToolUsage(web_fetch_requests=2, web_search_requests=3),
)

assert update_total_usage(first, total) is True
assert update_total_usage(second, total) is True
assert total.cache_creation == CacheCreation(
ephemeral_1h_input_tokens=20, ephemeral_5m_input_tokens=40
)
assert total.server_tool_use == ServerToolUsage(
web_fetch_requests=4, web_search_requests=6
)
assert second.cache_creation == total.cache_creation
assert second.server_tool_use == total.server_tool_use
85 changes: 85 additions & 0 deletions tests/coverage/test_core_patch_retry_coverage.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,6 +128,91 @@ def test_openai_usage_adds_token_details_and_copies_totals_to_response() -> None
assert response_prompt_details.cached_tokens == 25


def test_openai_usage_accumulates_sdk_extra_numeric_fields() -> None:
total_usage = CompletionUsage.model_validate(
{
"completion_tokens": 1,
"prompt_tokens": 2,
"total_tokens": 3,
"future_tokens": 4,
"future_cost": 1.5,
}
)
response_usage = CompletionUsage.model_validate(
{
"completion_tokens": 5,
"prompt_tokens": 6,
"total_tokens": 11,
"future_tokens": 7,
"future_cost": 2.25,
}
)
response = _completion(1)
response.usage = response_usage

update_total_usage(response, total_usage)

assert total_usage.model_extra is not None
assert response_usage.model_extra is not None
assert total_usage.model_extra["future_tokens"] == 11
assert total_usage.model_extra["future_cost"] == 3.75
assert response_usage.model_extra["future_tokens"] == 11
assert response_usage.model_extra["future_cost"] == 3.75


def test_usage_carries_existing_non_numeric_fields_to_response() -> None:
total_usage = CompletionUsage.model_validate(
{
"completion_tokens": 1,
"prompt_tokens": 2,
"total_tokens": 3,
"service_tier": "scale",
}
)
response_usage = CompletionUsage.model_validate(
{
"completion_tokens": 5,
"prompt_tokens": 6,
"total_tokens": 11,
}
)
response = _completion(1)
response.usage = response_usage

update_total_usage(response, total_usage)

assert response_usage.model_extra is not None
assert response_usage.model_extra["service_tier"] == "scale"


def test_usage_response_non_numeric_fields_still_win() -> None:
total_usage = CompletionUsage.model_validate(
{
"completion_tokens": 1,
"prompt_tokens": 2,
"total_tokens": 3,
"service_tier": "scale",
}
)
response_usage = CompletionUsage.model_validate(
{
"completion_tokens": 5,
"prompt_tokens": 6,
"total_tokens": 11,
"service_tier": "default",
}
)
response = _completion(1)
response.usage = response_usage

update_total_usage(response, total_usage)

assert total_usage.model_extra is not None
assert total_usage.model_extra["service_tier"] == "default"
assert response_usage.model_extra is not None
assert response_usage.model_extra["service_tier"] == "default"


def test_patch_requires_a_target_and_supports_a_create_callable() -> None:
patch_without_target = cast(Callable[[], object], patch)
with pytest.raises(ValueError, match="Either client or create must be provided"):
Expand Down