fix: Cost calculation enhancement (#207)

This commit is contained in:
Lifei Zhou
2024-10-31 13:58:51 +11:00
committed by GitHub
parent 0ab1966e93
commit aa324ce507
6 changed files with 89 additions and 27 deletions
+4 -1
View File
@@ -1,3 +1,4 @@
from datetime import datetime
import os
from typing import Union
from unittest.mock import MagicMock, mock_open, patch
@@ -159,13 +160,15 @@ def test_process_first_message_return_last_exchange_message(create_session_with_
def test_log_log_cost(create_session_with_mock_configs):
session = create_session_with_mock_configs()
mock_logger = MagicMock()
start_time = datetime(2024, 10, 20, 1, 2, 3)
end_time = datetime(2024, 10, 21, 2, 3, 4)
cost_message = "You have used 100 tokens"
with (
patch("exchange.Exchange.get_token_usage", return_value={}),
patch("goose.cli.session.get_total_cost_message", return_value=cost_message),
patch("goose.cli.session.get_logger", return_value=mock_logger),
):
session._log_cost()
session._log_cost(start_time, end_time)
mock_logger.info.assert_called_once_with(cost_message)
+42 -13
View File
@@ -1,7 +1,27 @@
from unittest.mock import patch
from datetime import datetime, timezone
from unittest.mock import MagicMock, patch
import pytest
from goose.utils._cost_calculator import _calculate_cost, get_total_cost_message
from exchange.providers.base import Usage
from goose.utils._cost_calculator import _calculate_cost, get_total_cost_message
SESSION_NAME = "test_session"
START_TIME = datetime(2024, 10, 20, 1, 2, 3, tzinfo=timezone.utc)
END_TIME = datetime(2024, 10, 21, 2, 3, 4, tzinfo=timezone.utc)
@pytest.fixture
def start_time():
mock_start_time = MagicMock(spec=datetime)
mock_start_time.astimezone.return_value = START_TIME
return mock_start_time
@pytest.fixture
def end_time():
mock_end_time = MagicMock(spec=datetime)
mock_end_time.astimezone.return_value = END_TIME
return mock_end_time
@pytest.fixture
@@ -16,32 +36,41 @@ def test_calculate_cost(mock_prices):
assert cost == 0.059
def test_get_total_cost_message(mock_prices):
def test_get_total_cost_message(mock_prices, start_time, end_time):
message = get_total_cost_message(
{
"gpt-4o": Usage(input_tokens=10000, output_tokens=600, total_tokens=10600),
"gpt-4o-mini": Usage(input_tokens=3000000, output_tokens=4000000, total_tokens=7000000),
}
},
SESSION_NAME,
start_time,
end_time,
)
expected_message = (
"Cost for model gpt-4o Usage(input_tokens=10000, output_tokens=600, total_tokens=10600): $0.06\n"
+ "Cost for model gpt-4o-mini Usage(input_tokens=3000000, output_tokens=4000000, total_tokens=7000000)"
+ ": $2.85\nTotal cost: $2.91"
"Session name: test_session | Cost for model gpt-4o Usage(input_tokens=10000, output_tokens=600,"
" total_tokens=10600): $0.06\n"
"Session name: test_session | Cost for model gpt-4o-mini Usage(input_tokens=3000000, output_tokens=4000000, "
"total_tokens=7000000): $2.85\n"
"2024-10-20T01:02:03+00:00 - 2024-10-21T02:03:04+00:00 | Session name: test_session | Total cost: $2.91"
)
assert message == expected_message
def test_get_total_cost_message_with_non_available_pricing(mock_prices):
def test_get_total_cost_message_with_non_available_pricing(mock_prices, start_time, end_time):
message = get_total_cost_message(
{
"non_pricing_model": Usage(input_tokens=10000, output_tokens=600, total_tokens=10600),
"gpt-4o-mini": Usage(input_tokens=3000000, output_tokens=4000000, total_tokens=7000000),
}
},
SESSION_NAME,
start_time,
end_time,
)
expected_message = (
"Cost for model non_pricing_model Usage(input_tokens=10000, output_tokens=600, total_tokens=10600): "
+ "Not available\n"
+ "Cost for model gpt-4o-mini Usage(input_tokens=3000000, output_tokens=4000000, total_tokens=7000000)"
+ ": $2.85\nTotal cost: $2.85"
"Session name: test_session | Cost for model non_pricing_model Usage(input_tokens=10000, output_tokens=600,"
" total_tokens=10600): Not available\n"
+ "Session name: test_session | Cost for model gpt-4o-mini Usage(input_tokens=3000000, output_tokens=4000000,"
" total_tokens=7000000): $2.85\n"
+ "2024-10-20T01:02:03+00:00 - 2024-10-21T02:03:04+00:00 | Session name: test_session | Total cost: $2.85"
)
assert message == expected_message