mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-04 02:01:48 +03:00
refactor(providers): define typed usage contract
This commit is contained in:
+11
-10
@@ -17,6 +17,7 @@ from nanobot.api.server import (
|
||||
create_app,
|
||||
handle_chat_completions,
|
||||
)
|
||||
from nanobot.providers.base import LLMUsage
|
||||
|
||||
try:
|
||||
from aiohttp.test_utils import TestClient, TestServer
|
||||
@@ -35,7 +36,7 @@ def _make_mock_agent(response_text: str = "mock response") -> MagicMock:
|
||||
agent = MagicMock()
|
||||
agent.process_direct = AsyncMock(return_value=response_text)
|
||||
agent.aclose = AsyncMock()
|
||||
agent._last_usage = {"prompt_tokens": 100, "completion_tokens": 50}
|
||||
agent._last_usage = LLMUsage.reported(input_tokens=100, output_tokens=50)
|
||||
return agent
|
||||
|
||||
|
||||
@@ -87,19 +88,19 @@ def test_chat_completion_response() -> None:
|
||||
|
||||
|
||||
def test_chat_completion_response_with_usage() -> None:
|
||||
usage = {"prompt_tokens": 150, "completion_tokens": 42}
|
||||
usage = LLMUsage.reported(input_tokens=150, output_tokens=42)
|
||||
result = _chat_completion_response("hello world", "test-model", usage)
|
||||
assert result["usage"]["prompt_tokens"] == 150
|
||||
assert result["usage"]["completion_tokens"] == 42
|
||||
assert result["usage"]["total_tokens"] == 192
|
||||
|
||||
|
||||
def test_chat_completion_response_preserves_provider_total_usage() -> None:
|
||||
usage = {"total_tokens": 77}
|
||||
def test_chat_completion_response_preserves_explicit_total_usage() -> None:
|
||||
usage = LLMUsage.reported(input_tokens=70, output_tokens=7, total_tokens=175)
|
||||
result = _chat_completion_response("hello world", "test-model", usage)
|
||||
assert result["usage"]["prompt_tokens"] == 0
|
||||
assert result["usage"]["completion_tokens"] == 0
|
||||
assert result["usage"]["total_tokens"] == 77
|
||||
assert result["usage"]["prompt_tokens"] == 70
|
||||
assert result["usage"]["completion_tokens"] == 7
|
||||
assert result["usage"]["total_tokens"] == 175
|
||||
|
||||
|
||||
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
|
||||
@@ -328,7 +329,7 @@ async def test_followup_requests_share_same_session_key(aiohttp_client) -> None:
|
||||
agent = MagicMock()
|
||||
agent.process_direct = fake_process
|
||||
agent.aclose = AsyncMock()
|
||||
agent._last_usage = {}
|
||||
agent._last_usage = None
|
||||
|
||||
app = create_app(agent, model_name="m", api_key=API_KEY)
|
||||
client = await aiohttp_client(app)
|
||||
@@ -367,7 +368,7 @@ async def test_fixed_session_requests_are_serialized(aiohttp_client) -> None:
|
||||
agent = MagicMock()
|
||||
agent.process_direct = slow_process
|
||||
agent.aclose = AsyncMock()
|
||||
agent._last_usage = {}
|
||||
agent._last_usage = None
|
||||
|
||||
app = create_app(agent, model_name="m", api_key=API_KEY)
|
||||
client = await aiohttp_client(app)
|
||||
@@ -484,7 +485,7 @@ async def test_empty_response_falls_back_without_retry(aiohttp_client) -> None:
|
||||
agent = MagicMock()
|
||||
agent.process_direct = always_empty
|
||||
agent.aclose = AsyncMock()
|
||||
agent._last_usage = {}
|
||||
agent._last_usage = None
|
||||
|
||||
app = create_app(agent, model_name="m", api_key=API_KEY)
|
||||
client = await aiohttp_client(app)
|
||||
|
||||
Reference in New Issue
Block a user