From ee100e40bc4d662f379d913b9403f30439de5c33 Mon Sep 17 00:00:00 2001 From: Rhonin Date: Sat, 15 Aug 2026 02:34:45 +0800 Subject: [PATCH 1/8] test: cover provider stats call paths --- tests/unit/test_cron_manager.py | 95 +++++++++++++++++++++- tests/unit/test_star_context.py | 138 ++++++++++++++++++++++++++++++++ 2 files changed, 232 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_cron_manager.py b/tests/unit/test_cron_manager.py index 47f97ef445..5b0430804f 100644 --- a/tests/unit/test_cron_manager.py +++ b/tests/unit/test_cron_manager.py @@ -1,17 +1,21 @@ """Tests for CronJobManager.""" from datetime import datetime, timedelta, timezone +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch from zoneinfo import ZoneInfo import pytest +from sqlmodel import select +from astrbot.core.agent.response import AgentStats from astrbot.core.cron.manager import ( CronJobManager, CronJobSchedulingError, _normalize_crontab_day_of_week, ) -from astrbot.core.db.po import CronJob +from astrbot.core.db.po import CronJob, ProviderStat +from astrbot.core.provider.entities import LLMResponse, TokenUsage @pytest.fixture @@ -646,6 +650,95 @@ async def fake_persist_agent_history(*args, **kwargs): assert config.provider_settings is provider_settings assert config.provider_settings["fallback_chat_models"] == ["fallback-provider"] + @pytest.mark.asyncio + async def test_woke_main_agent_persists_one_aggregated_provider_stat( + self, + temp_db, + ): + manager = CronJobManager(temp_db) + ctx = MagicMock() + ctx.get_config.return_value = { + "admins_id": [], + "provider_settings": {}, + } + ctx.conversation_manager = MagicMock() + manager.ctx = ctx + + conv = MagicMock() + conv.cid = "conv-cron" + conv.history = "[]" + final_response = LLMResponse( + role="assistant", + completion_text="done", + usage=TokenUsage(input_other=9, input_cached=2, output=5), + ) + provider = SimpleNamespace( + provider_config={"id": "provider-cron"}, + meta=lambda: SimpleNamespace(id="provider-cron", type="test"), + get_model=lambda: "cron-model", + ) + + class FakeRunner: + def __init__(self) -> None: + self.provider = provider + self.stats = AgentStats( + token_usage=TokenUsage( + input_other=20, + input_cached=4, + output=10, + ), + start_time=200.0, + end_time=212.0, + time_to_first_token=0.7, + ) + + async def step_until_done(self, max_steps): + if False: + yield None + + def get_final_llm_resp(self): + return final_response + + def was_aborted(self) -> bool: + return False + + async def fake_build_main_agent(*, event, plugin_context, config, req): + return MagicMock(agent_runner=FakeRunner()) + + with ( + patch( + "astrbot.core.astr_main_agent._get_session_conv", + AsyncMock(return_value=conv), + ), + patch( + "astrbot.core.astr_main_agent.build_main_agent", + side_effect=fake_build_main_agent, + ), + patch( + "astrbot.core.cron.manager.persist_agent_history", + new=AsyncMock(), + ), + ): + await manager._woke_main_agent( + message="scheduled task", + session_str="test:FriendMessage:user123", + extras={"cron_job": {"id": "job-1"}, "cron_payload": {}}, + ) + + async with temp_db.get_db() as session: + result = await session.execute(select(ProviderStat)) + records = result.scalars().all() + + assert len(records) == 1 + record = records[0] + assert record.agent_type == "internal" + assert record.conversation_id == "conv-cron" + assert record.provider_id == "provider-cron" + assert record.provider_model == "cron-model" + assert record.token_input_other == 20 + assert record.token_input_cached == 4 + assert record.token_output == 10 + class TestGetNextRunTime: """Tests for _get_next_run_time method.""" diff --git a/tests/unit/test_star_context.py b/tests/unit/test_star_context.py index 0979aaa0bf..c07e5ffe1a 100644 --- a/tests/unit/test_star_context.py +++ b/tests/unit/test_star_context.py @@ -1,9 +1,15 @@ from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest +from sqlmodel import select +from astrbot.core.agent.response import AgentStats from astrbot.core.agent.tool import FunctionTool +from astrbot.core.db.po import ProviderStat +from astrbot.core.provider.entities import LLMResponse, ProviderMeta, TokenUsage from astrbot.core.provider.func_tool_manager import FunctionToolManager +from astrbot.core.provider.provider import Provider from astrbot.core.star.context import Context from astrbot.core.star.star import StarMetadata, star_registry @@ -34,6 +40,41 @@ def make_tool(name: str, module_path: str) -> FunctionTool: return tool +class StatsProvider(Provider): + def __init__(self) -> None: + super().__init__({"id": "provider-1", "type": "test"}, {}) + self.set_model("test-model") + + def get_current_key(self) -> str: + return "" + + def set_key(self, key: str) -> None: + return None + + async def get_models(self) -> list[str]: + return [self.get_model()] + + def meta(self) -> ProviderMeta: + return ProviderMeta( + id="provider-1", + model=self.get_model(), + type="test", + ) + + async def text_chat(self, **kwargs) -> LLMResponse: + return LLMResponse( + role="assistant", + completion_text="ok", + usage=TokenUsage(input_other=5, input_cached=2, output=3), + ) + + +async def get_provider_stats(temp_db) -> list[ProviderStat]: + async with temp_db.get_db() as session: + result = await session.execute(select(ProviderStat)) + return list(result.scalars().all()) + + def test_add_llm_tools_resolves_subdirectory_plugin_without_name_prefix(): star_registry.append( StarMetadata( @@ -104,3 +145,100 @@ def test_add_llm_tools_handles_empty_tool_module_path(): context.add_llm_tools(tool) assert tool.handler_module_path == "" + + +@pytest.mark.asyncio +async def test_llm_generate_persists_one_provider_stat(temp_db): + provider = StatsProvider() + context = Context.__new__(Context) + context._db = temp_db + context.provider_manager = SimpleNamespace( + get_provider_by_id=AsyncMock(return_value=provider), + ) + + response = await context.llm_generate( + chat_provider_id="provider-1", + prompt="test", + session_id="session-1", + ) + + records = await get_provider_stats(temp_db) + assert response.completion_text == "ok" + assert len(records) == 1 + record = records[0] + assert record.agent_type == "internal" + assert record.status == "completed" + assert record.umo == "provider:provider-1:session-1" + assert record.provider_id == "provider-1" + assert record.provider_model == "test-model" + assert record.token_input_other == 5 + assert record.token_input_cached == 2 + assert record.token_output == 3 + assert record.end_time >= record.start_time > 0 + + +@pytest.mark.asyncio +async def test_tool_loop_agent_persists_one_aggregated_provider_stat( + temp_db, + monkeypatch: pytest.MonkeyPatch, +): + provider = StatsProvider() + final_response = LLMResponse( + role="assistant", + completion_text="done", + usage=TokenUsage(input_other=8, input_cached=1, output=4), + ) + + class FakeRunner: + def __init__(self) -> None: + self.provider = provider + self.stats = AgentStats( + token_usage=TokenUsage(input_other=12, input_cached=3, output=7), + start_time=100.0, + end_time=106.0, + time_to_first_token=0.4, + ) + + async def reset(self, **kwargs) -> None: + return None + + async def step_until_done(self, max_steps): + if False: + yield None + + def get_final_llm_resp(self) -> LLMResponse: + return final_response + + def was_aborted(self) -> bool: + return False + + monkeypatch.setattr("astrbot.core.star.context.ToolLoopAgentRunner", FakeRunner) + + context = Context.__new__(Context) + context._db = temp_db + context.provider_manager = SimpleNamespace( + get_provider_by_id=AsyncMock(return_value=provider), + ) + event = SimpleNamespace( + unified_msg_origin="webchat:FriendMessage:session-42", + ) + + response = await context.tool_loop_agent( + event=event, + chat_provider_id="provider-1", + prompt="test", + agent_context=SimpleNamespace(), + ) + + records = await get_provider_stats(temp_db) + assert response is final_response + assert len(records) == 1 + record = records[0] + assert record.agent_type == "internal" + assert record.umo == "webchat:FriendMessage:session-42" + assert record.provider_id == "provider-1" + assert record.token_input_other == 12 + assert record.token_input_cached == 3 + assert record.token_output == 7 + assert record.start_time == 100.0 + assert record.end_time == 106.0 From 4b93c4ee0e52db214eccd4477c12495a0542fd9e Mon Sep 17 00:00:00 2001 From: Rhonin Date: Sat, 15 Aug 2026 02:38:29 +0800 Subject: [PATCH 2/8] fix: record provider stats for detached calls --- astrbot/core/cron/manager.py | 8 ++ .../method/agent_sub_stages/internal.py | 42 ++------ astrbot/core/provider/stats.py | 98 +++++++++++++++++++ astrbot/core/star/context.py | 44 +++++++-- 4 files changed, 149 insertions(+), 43 deletions(-) create mode 100644 astrbot/core/provider/stats.py diff --git a/astrbot/core/cron/manager.py b/astrbot/core/cron/manager.py index b5a0e7c3e4..95e3182d15 100644 --- a/astrbot/core/cron/manager.py +++ b/astrbot/core/cron/manager.py @@ -18,6 +18,7 @@ from astrbot.core.platform.message_session import MessageSession from astrbot.core.platform.message_type import MessageType from astrbot.core.provider.entites import ProviderRequest +from astrbot.core.provider.stats import record_agent_runner_stats from astrbot.core.utils.history_saver import persist_agent_history if TYPE_CHECKING: @@ -492,6 +493,13 @@ async def _woke_main_agent( # agent will send message to user via using tools pass llm_resp = runner.get_final_llm_resp() + await record_agent_runner_stats( + self.db, + umo=cron_event.unified_msg_origin, + request=req, + agent_runner=runner, + final_response=llm_resp, + ) cron_meta = extras.get("cron_job", {}) if extras else {} summary_note = ( f"[CronJob] {cron_meta.get('name') or cron_meta.get('id', 'unknown')}: {cron_meta.get('description', '')} " diff --git a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py index 40e0e99a50..b58e3a7abc 100644 --- a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py +++ b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py @@ -34,6 +34,7 @@ LLMResponse, ProviderRequest, ) +from astrbot.core.provider.stats import record_agent_runner_stats from astrbot.core.star.star_handler import EventType from astrbot.core.utils.metrics import Metric from astrbot.core.utils.session_lock import session_lock_manager @@ -550,37 +551,10 @@ async def _record_internal_agent_stats( final_resp: LLMResponse | None, ) -> None: """Persist internal agent stats without affecting the user response flow.""" - if agent_runner is None: - return - - provider = agent_runner.provider - stats = agent_runner.stats - if provider is None or stats is None: - return - - try: - provider_config = getattr(provider, "provider_config", {}) or {} - conversation_id = ( - req.conversation.cid - if req is not None and req.conversation is not None - else None - ) - - if agent_runner.was_aborted(): - status = "aborted" - elif final_resp is not None and final_resp.role == "err": - status = "error" - else: - status = "completed" - - await db_helper.insert_provider_stat( - umo=event.unified_msg_origin, - conversation_id=conversation_id, - provider_id=provider_config.get("id", "") or provider.meta().id, - provider_model=provider.get_model(), - status=status, - stats=stats.to_dict(), - agent_type="internal", - ) - except Exception as e: - logger.warning("Persist provider stats failed: %s", e, exc_info=True) + await record_agent_runner_stats( + db_helper, + umo=event.unified_msg_origin, + request=req, + agent_runner=agent_runner, + final_response=final_resp, + ) diff --git a/astrbot/core/provider/stats.py b/astrbot/core/provider/stats.py new file mode 100644 index 0000000000..a616a7aeb9 --- /dev/null +++ b/astrbot/core/provider/stats.py @@ -0,0 +1,98 @@ +from __future__ import annotations + +from typing import Any + +from astrbot import logger +from astrbot.core.db import BaseDatabase +from astrbot.core.provider.entities import LLMResponse, ProviderRequest, TokenUsage + + +def _provider_id(provider: Any) -> str: + provider_config = getattr(provider, "provider_config", {}) or {} + return provider_config.get("id", "") or provider.meta().id + + +def _response_status(response: LLMResponse | None) -> str: + if response is None or response.role == "err": + return "error" + return "completed" + + +def _runner_status(response: LLMResponse | None, aborted: bool) -> str: + if aborted: + return "aborted" + if response is not None and response.role == "err": + return "error" + return "completed" + + +async def record_agent_runner_stats( + db: BaseDatabase, + *, + umo: str, + request: ProviderRequest | None, + agent_runner: Any, + final_response: LLMResponse | None, + agent_type: str = "internal", +) -> None: + """Persist aggregate agent runner stats without affecting its response.""" + if agent_runner is None: + return + + provider = getattr(agent_runner, "provider", None) + stats = getattr(agent_runner, "stats", None) + if provider is None or stats is None: + return + + try: + conversation_id = ( + request.conversation.cid + if request is not None and request.conversation is not None + else None + ) + await db.insert_provider_stat( + umo=umo, + conversation_id=conversation_id, + provider_id=_provider_id(provider), + provider_model=provider.get_model(), + status=_runner_status( + final_response, + agent_runner.was_aborted(), + ), + stats=stats.to_dict(), + agent_type=agent_type, + ) + except Exception as exc: # noqa: BLE001 + logger.warning("Persist provider stats failed: %s", exc, exc_info=True) + + +async def record_llm_response_stats( + db: BaseDatabase, + *, + umo: str, + provider: Any, + response: LLMResponse | None, + start_time: float, + end_time: float, + conversation_id: str | None = None, + agent_type: str = "internal", +) -> None: + """Persist stats for one direct provider request.""" + try: + usage = response.usage if response and response.usage else TokenUsage() + await db.insert_provider_stat( + umo=umo, + conversation_id=conversation_id, + provider_id=_provider_id(provider), + provider_model=provider.get_model(), + status=_response_status(response), + stats={ + "token_usage": usage.__dict__.copy(), + "start_time": start_time, + "end_time": end_time, + "time_to_first_token": 0.0, + }, + agent_type=agent_type, + ) + except Exception as exc: # noqa: BLE001 + logger.warning("Persist provider stats failed: %s", exc, exc_info=True) diff --git a/astrbot/core/star/context.py b/astrbot/core/star/context.py index b4f6e61c48..25972a502f 100644 --- a/astrbot/core/star/context.py +++ b/astrbot/core/star/context.py @@ -1,6 +1,7 @@ from __future__ import annotations import logging +import time from asyncio import Queue from collections.abc import Awaitable, Callable from typing import TYPE_CHECKING, Any, Protocol @@ -32,6 +33,10 @@ STTProvider, TTSProvider, ) +from astrbot.core.provider.stats import ( + record_agent_runner_stats, + record_llm_response_stats, +) from astrbot.core.star.filter.platform_adapter_type import ( ADAPTER_NAME_2_TYPE, PlatformAdapterType, @@ -201,15 +206,29 @@ async def llm_generate( prov = await self.provider_manager.get_provider_by_id(chat_provider_id) if not prov or not isinstance(prov, Provider): raise ProviderNotFoundError(f"Provider {chat_provider_id} not found") - llm_resp = await prov.text_chat( - prompt=prompt, - image_urls=image_urls, - audio_urls=audio_urls, - func_tool=tools, - contexts=contexts, - system_prompt=system_prompt, - **kwargs, - ) + start_time = time.time() + llm_resp = None + try: + llm_resp = await prov.text_chat( + prompt=prompt, + image_urls=image_urls, + audio_urls=audio_urls, + func_tool=tools, + contexts=contexts, + system_prompt=system_prompt, + **kwargs, + ) + finally: + session_id = kwargs.get("session_id") or "sdk" + await record_llm_response_stats( + self._db, + umo=f"provider:{prov.meta().id}:{session_id}", + provider=prov, + response=llm_resp, + start_time=start_time, + end_time=time.time(), + conversation_id=kwargs.get("conversation_id"), + ) return llm_resp async def tool_loop_agent( @@ -322,6 +341,13 @@ async def tool_loop_agent( async for _ in agent_runner.step_until_done(max_steps): pass llm_resp = agent_runner.get_final_llm_resp() + await record_agent_runner_stats( + self._db, + umo=event.unified_msg_origin, + request=request, + agent_runner=agent_runner, + final_response=llm_resp, + ) if not llm_resp: raise Exception("Agent did not produce a final LLM response") return llm_resp From ef3ed0f28dd583b115457b727f203e64b70ed512 Mon Sep 17 00:00:00 2001 From: Rhonin Date: Sat, 15 Aug 2026 03:15:04 +0800 Subject: [PATCH 3/8] test: cover failed detached provider stats --- tests/unit/test_cron_manager.py | 72 +++++++++++++++++++++++++++++++++ tests/unit/test_star_context.py | 57 +++++++++++++++++++++++++- tests/unit/test_stat_service.py | 34 ++++++++++++++++ 3 files changed, 162 insertions(+), 1 deletion(-) create mode 100644 tests/unit/test_stat_service.py diff --git a/tests/unit/test_cron_manager.py b/tests/unit/test_cron_manager.py index 5b0430804f..7b31976df9 100644 --- a/tests/unit/test_cron_manager.py +++ b/tests/unit/test_cron_manager.py @@ -739,6 +739,78 @@ async def fake_build_main_agent(*, event, plugin_context, config, req): assert record.token_input_cached == 4 assert record.token_output == 10 + @pytest.mark.asyncio + async def test_woke_main_agent_persists_failed_provider_stat(self, temp_db): + manager = CronJobManager(temp_db) + ctx = MagicMock() + ctx.get_config.return_value = { + "admins_id": [], + "provider_settings": {}, + } + ctx.conversation_manager = MagicMock() + manager.ctx = ctx + + conv = MagicMock() + conv.cid = "conv-cron-failed" + conv.history = "[]" + provider = SimpleNamespace( + provider_config={"id": "provider-cron"}, + meta=lambda: SimpleNamespace(id="provider-cron", type="test"), + get_model=lambda: "cron-model", + ) + + class FakeRunner: + def __init__(self) -> None: + self.provider = provider + self.stats = AgentStats( + token_usage=TokenUsage(input_other=6, output=3), + start_time=200.0, + end_time=201.0, + ) + + async def step_until_done(self, max_steps): + raise RuntimeError("cron provider failed") + yield + + def get_final_llm_resp(self): + return None + + def was_aborted(self) -> bool: + return False + + async def fake_build_main_agent(*, event, plugin_context, config, req): + return MagicMock(agent_runner=FakeRunner()) + + with ( + patch( + "astrbot.core.astr_main_agent._get_session_conv", + AsyncMock(return_value=conv), + ), + patch( + "astrbot.core.astr_main_agent.build_main_agent", + side_effect=fake_build_main_agent, + ), + patch( + "astrbot.core.cron.manager.persist_agent_history", + new=AsyncMock(), + ), + ): + with pytest.raises(RuntimeError, match="cron provider failed"): + await manager._woke_main_agent( + message="scheduled task", + session_str="test:FriendMessage:user123", + extras={"cron_job": {"id": "job-1"}, "cron_payload": {}}, + ) + + async with temp_db.get_db() as session: + result = await session.execute(select(ProviderStat)) + records = result.scalars().all() + + assert len(records) == 1 + assert records[0].status == "error" + assert records[0].token_input_other == 6 + assert records[0].token_output == 3 + class TestGetNextRunTime: """Tests for _get_next_run_time method.""" diff --git a/tests/unit/test_star_context.py b/tests/unit/test_star_context.py index c07e5ffe1a..9a7862d3a2 100644 --- a/tests/unit/test_star_context.py +++ b/tests/unit/test_star_context.py @@ -166,7 +166,7 @@ async def test_llm_generate_persists_one_provider_stat(temp_db): assert response.completion_text == "ok" assert len(records) == 1 record = records[0] - assert record.agent_type == "internal" + assert record.agent_type == "provider" assert record.status == "completed" assert record.umo == "provider:provider-1:session-1" assert record.provider_id == "provider-1" @@ -242,3 +242,58 @@ def was_aborted(self) -> bool: assert record.token_output == 7 assert record.start_time == 100.0 assert record.end_time == 106.0 + + +@pytest.mark.asyncio +async def test_tool_loop_agent_persists_failed_provider_stat( + temp_db, + monkeypatch: pytest.MonkeyPatch, +): + provider = StatsProvider() + + class FakeRunner: + def __init__(self) -> None: + self.provider = provider + self.stats = AgentStats( + token_usage=TokenUsage(input_other=4, output=2), + start_time=100.0, + end_time=101.0, + ) + + async def reset(self, **kwargs) -> None: + return None + + async def step_until_done(self, max_steps): + raise RuntimeError("provider failed") + yield + + def get_final_llm_resp(self): + return None + + def was_aborted(self) -> bool: + return False + + monkeypatch.setattr("astrbot.core.star.context.ToolLoopAgentRunner", FakeRunner) + + context = Context.__new__(Context) + context._db = temp_db + context.provider_manager = SimpleNamespace( + get_provider_by_id=AsyncMock(return_value=provider), + ) + event = SimpleNamespace( + unified_msg_origin="webchat:FriendMessage:failed-session", + ) + + with pytest.raises(RuntimeError, match="provider failed"): + await context.tool_loop_agent( + event=event, + chat_provider_id="provider-1", + prompt="test", + agent_context=SimpleNamespace(), + ) + + records = await get_provider_stats(temp_db) + assert len(records) == 1 + assert records[0].status == "error" + assert records[0].token_input_other == 4 + assert records[0].token_output == 2 diff --git a/tests/unit/test_stat_service.py b/tests/unit/test_stat_service.py new file mode 100644 index 0000000000..02813cbd4a --- /dev/null +++ b/tests/unit/test_stat_service.py @@ -0,0 +1,34 @@ +from types import SimpleNamespace + +import pytest + +from astrbot.dashboard.services.stat_service import StatService + + +@pytest.mark.asyncio +async def test_provider_token_stats_include_detached_provider_calls(temp_db): + for agent_type, provider_id, status, output in ( + ("internal", "agent", "completed", 3), + ("provider", "sdk", "completed", 5), + ("internal", "aborted", "aborted", 7), + ("third_party", "excluded", "completed", 100), + ): + await temp_db.insert_provider_stat( + umo=f"test:{provider_id}", + provider_id=provider_id, + provider_model=f"{provider_id}-model", + status=status, + stats={ + "token_usage": {"output": output}, + "start_time": 1.0, + "end_time": 2.0, + }, + agent_type=agent_type, + ) + + service = StatService(temp_db, SimpleNamespace(), {}) + stats = await service.get_provider_token_stats(1) + + assert stats["range_total_calls"] == 3 + assert stats["range_total_tokens"] == 15 + assert stats["range_success_rate"] == pytest.approx(2 / 3) From 5a43deb3b912efd6a8a89928730868c91d0b62d7 Mon Sep 17 00:00:00 2001 From: Rhonin Date: Sat, 15 Aug 2026 03:22:44 +0800 Subject: [PATCH 4/8] fix: preserve detached provider stats on failures --- astrbot/core/cron/manager.py | 25 ++++++++++++--------- astrbot/core/provider/stats.py | 2 +- astrbot/core/star/context.py | 26 +++++++++++++--------- astrbot/dashboard/services/stat_service.py | 4 ++-- 4 files changed, 32 insertions(+), 25 deletions(-) diff --git a/astrbot/core/cron/manager.py b/astrbot/core/cron/manager.py index 95e3182d15..0fc9543891 100644 --- a/astrbot/core/cron/manager.py +++ b/astrbot/core/cron/manager.py @@ -489,17 +489,20 @@ async def _woke_main_agent( return runner = result.agent_runner - async for _ in runner.step_until_done(30): - # agent will send message to user via using tools - pass - llm_resp = runner.get_final_llm_resp() - await record_agent_runner_stats( - self.db, - umo=cron_event.unified_msg_origin, - request=req, - agent_runner=runner, - final_response=llm_resp, - ) + llm_resp = None + try: + async for _ in runner.step_until_done(30): + # agent will send message to user via using tools + pass + llm_resp = runner.get_final_llm_resp() + finally: + await record_agent_runner_stats( + self.db, + umo=cron_event.unified_msg_origin, + request=req, + agent_runner=runner, + final_response=llm_resp, + ) cron_meta = extras.get("cron_job", {}) if extras else {} summary_note = ( f"[CronJob] {cron_meta.get('name') or cron_meta.get('id', 'unknown')}: {cron_meta.get('description', '')} " diff --git a/astrbot/core/provider/stats.py b/astrbot/core/provider/stats.py index a616a7aeb9..26c61e19c1 100644 --- a/astrbot/core/provider/stats.py +++ b/astrbot/core/provider/stats.py @@ -21,7 +21,7 @@ def _response_status(response: LLMResponse | None) -> str: def _runner_status(response: LLMResponse | None, aborted: bool) -> str: if aborted: return "aborted" - if response is not None and response.role == "err": + if response is None or response.role == "err": return "error" return "completed" diff --git a/astrbot/core/star/context.py b/astrbot/core/star/context.py index 25972a502f..f4b9679b7b 100644 --- a/astrbot/core/star/context.py +++ b/astrbot/core/star/context.py @@ -222,12 +222,13 @@ async def llm_generate( session_id = kwargs.get("session_id") or "sdk" await record_llm_response_stats( self._db, - umo=f"provider:{prov.meta().id}:{session_id}", + umo=f"provider:{chat_provider_id}:{session_id}", provider=prov, response=llm_resp, start_time=start_time, end_time=time.time(), conversation_id=kwargs.get("conversation_id"), + agent_type="provider", ) return llm_resp @@ -338,16 +339,19 @@ async def tool_loop_agent( streaming=streaming, **other_kwargs, ) - async for _ in agent_runner.step_until_done(max_steps): - pass - llm_resp = agent_runner.get_final_llm_resp() - await record_agent_runner_stats( - self._db, - umo=event.unified_msg_origin, - request=request, - agent_runner=agent_runner, - final_response=llm_resp, - ) + llm_resp = None + try: + async for _ in agent_runner.step_until_done(max_steps): + pass + llm_resp = agent_runner.get_final_llm_resp() + finally: + await record_agent_runner_stats( + self._db, + umo=event.unified_msg_origin, + request=request, + agent_runner=agent_runner, + final_response=llm_resp, + ) if not llm_resp: raise Exception("Agent did not produce a final LLM response") return llm_resp diff --git a/astrbot/dashboard/services/stat_service.py b/astrbot/dashboard/services/stat_service.py index 061702d659..a5bf1932f9 100644 --- a/astrbot/dashboard/services/stat_service.py +++ b/astrbot/dashboard/services/stat_service.py @@ -299,7 +299,7 @@ async def get_provider_token_stats(self, days: int) -> dict: result = await session.execute( select(ProviderStat) .where( - ProviderStat.agent_type == "internal", + col(ProviderStat.agent_type).in_(("internal", "provider")), ProviderStat.created_at >= query_start_utc, ) .order_by(col(ProviderStat.created_at).asc()) @@ -353,7 +353,7 @@ async def get_provider_token_stats(self, days: int) -> dict: total_by_bucket[bucket_ts] += token_total range_total_tokens += token_total range_total_calls += 1 - if record.status != "error": + if record.status == "completed": range_success_calls += 1 if record.time_to_first_token > 0: range_ttft_total_ms += record.time_to_first_token * 1000 From 5149f2e08407d42a4500adb3fbaa78fd1254a123 Mon Sep 17 00:00:00 2001 From: Rhonin Date: Sat, 15 Aug 2026 04:11:00 +0800 Subject: [PATCH 5/8] test: cover cancellation and fallback attribution --- tests/test_tool_loop_agent_runner.py | 49 +++++++++++++ tests/unit/test_provider_stats.py | 105 +++++++++++++++++++++++++++ 2 files changed, 154 insertions(+) diff --git a/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index 1e679de4aa..e316bd9b12 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -183,6 +183,18 @@ async def text_chat(self, **kwargs) -> LLMResponse: raise RuntimeError("primary provider failed") +class MockUsageFailingProvider(MockProvider): + async def text_chat(self, **kwargs) -> LLMResponse: + self.call_count += 1 + error = RuntimeError("primary response parsing failed") + error._astrbot_token_usage = TokenUsage( # type: ignore[attr-defined] + input_other=8, + input_cached=4, + output=6, + ) + raise error + + class MockErrProvider(MockProvider): async def text_chat(self, **kwargs) -> LLMResponse: self.call_count += 1 @@ -1213,6 +1225,43 @@ async def test_fallback_provider_used_when_primary_raises( assert fallback_provider.call_count == 1 +@pytest.mark.asyncio +async def test_fallback_tracks_failed_primary_usage_by_provider( + runner, + provider_request, + mock_tool_executor, + mock_hooks, +): + primary_provider = MockUsageFailingProvider() + primary_provider.provider_config["id"] = "primary" + fallback_provider = MockProvider() + fallback_provider.provider_config["id"] = "fallback" + fallback_provider.should_call_tools = False + + await runner.reset( + provider=primary_provider, + request=provider_request, + run_context=ContextWrapper(context=None), + tool_executor=mock_tool_executor, + agent_hooks=mock_hooks, + streaming=False, + fallback_providers=[fallback_provider], + ) + + async for _ in runner.step_until_done(5): + pass + + assert runner.stats.token_usage == TokenUsage( + input_other=18, + input_cached=4, + output=11, + ) + assert len(runner.provider_stat_segments) == 1 + segment = runner.provider_stat_segments[0] + assert segment.provider is primary_provider + assert segment.usage == TokenUsage(input_other=8, input_cached=4, output=6) + + @pytest.mark.asyncio async def test_fallback_provider_used_when_primary_returns_err( runner, provider_request, mock_tool_executor, mock_hooks diff --git a/tests/unit/test_provider_stats.py b/tests/unit/test_provider_stats.py index c306c3f820..cb0a4c4b45 100644 --- a/tests/unit/test_provider_stats.py +++ b/tests/unit/test_provider_stats.py @@ -1,4 +1,6 @@ +import asyncio from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest from sqlmodel import select @@ -7,6 +9,7 @@ from astrbot.core.db.po import ProviderStat from astrbot.core.pipeline.process_stage.method.agent_sub_stages import internal from astrbot.core.provider.entities import ProviderRequest, TokenUsage +from astrbot.core.provider.stats import ProviderStatSegment @pytest.mark.asyncio @@ -63,3 +66,105 @@ async def test_record_internal_agent_stats_persists_provider_stat( assert record.start_time == 100.0 assert record.end_time == 108.5 assert record.time_to_first_token == 0.6 + + +@pytest.mark.asyncio +async def test_record_internal_agent_stats_splits_failed_fallback_provider_usage( + temp_db, + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr(internal, "db_helper", temp_db) + primary = SimpleNamespace( + provider_config={"id": "primary"}, + meta=lambda: SimpleNamespace(id="primary", type="openai"), + get_model=lambda: "primary-model", + ) + fallback = SimpleNamespace( + provider_config={"id": "fallback"}, + meta=lambda: SimpleNamespace(id="fallback", type="openai"), + get_model=lambda: "fallback-model", + ) + runner = SimpleNamespace( + provider=fallback, + stats=AgentStats( + token_usage=TokenUsage(input_other=18, input_cached=4, output=11), + start_time=100.0, + end_time=108.0, + time_to_first_token=5.0, + ), + provider_stat_segments=[ + ProviderStatSegment( + provider=primary, + usage=TokenUsage(input_other=8, input_cached=4, output=6), + start_time=100.0, + end_time=103.0, + ) + ], + was_aborted=lambda: False, + ) + + await internal._record_internal_agent_stats( + SimpleNamespace(unified_msg_origin="test:session"), + ProviderRequest(conversation=SimpleNamespace(cid="conv-1")), + runner, + SimpleNamespace(role="assistant"), + ) + + async with temp_db.get_db() as session: + result = await session.execute(select(ProviderStat)) + records = sorted(result.scalars().all(), key=lambda item: item.provider_id) + + assert len(records) == 2 + fallback_record, primary_record = records + assert fallback_record.provider_id == "fallback" + assert fallback_record.status == "completed" + assert fallback_record.token_input_other == 10 + assert fallback_record.token_input_cached == 0 + assert fallback_record.token_output == 5 + assert primary_record.provider_id == "primary" + assert primary_record.status == "error" + assert primary_record.token_input_other == 8 + assert primary_record.token_input_cached == 4 + assert primary_record.token_output == 6 + + +@pytest.mark.asyncio +async def test_cancelled_agent_finally_schedules_stats_once( + monkeypatch: pytest.MonkeyPatch, +): + writer = AsyncMock() + monkeypatch.setattr(internal, "_record_internal_agent_stats", writer) + event = SimpleNamespace(unified_msg_origin="test:cancelled") + request = ProviderRequest() + runner = SimpleNamespace(get_final_llm_resp=lambda: None) + started = asyncio.Event() + + async def cancelled_run() -> None: + scheduled = False + try: + started.set() + await asyncio.Event().wait() + finally: + scheduled = internal._schedule_internal_agent_stats( + scheduled, + event, + request, + runner, + runner.get_final_llm_resp(), + ) + internal._schedule_internal_agent_stats( + scheduled, + event, + request, + runner, + runner.get_final_llm_resp(), + ) + + task = asyncio.create_task(cancelled_run()) + await started.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + await asyncio.sleep(0) + + writer.assert_awaited_once_with(event, request, runner, None) From b711bacc51540b264963518396fe2a73ec31fcdf Mon Sep 17 00:00:00 2001 From: Rhonin Date: Sat, 15 Aug 2026 04:14:01 +0800 Subject: [PATCH 6/8] fix: preserve provider attribution on fallback --- .../agent/runners/tool_loop_agent_runner.py | 26 +++++++++ .../method/agent_sub_stages/internal.py | 38 ++++++++++--- astrbot/core/provider/stats.py | 53 ++++++++++++++++++- 3 files changed, 109 insertions(+), 8 deletions(-) diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index 8c91adbbfd..5b361408ff 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -39,6 +39,7 @@ from astrbot.core.provider.entities import ( LLMResponse, ProviderRequest, + TokenUsage, ToolCallsResult, ) from astrbot.core.provider.modalities import ( @@ -46,6 +47,7 @@ sanitize_contexts_by_modalities, ) from astrbot.core.provider.provider import Provider +from astrbot.core.provider.stats import ProviderStatSegment from ..context.compressor import ContextCompressor from ..context.config import ContextConfig @@ -324,6 +326,7 @@ async def reset( self.stats = AgentStats() self.stats.start_time = time.time() + self.provider_stat_segments: list[ProviderStatSegment] = [] def _read_tool_hint(self) -> str: if self.read_tool is not None: @@ -551,6 +554,7 @@ async def _iter_llm_responses_with_fallback( candidate_id, ) self.provider = candidate + candidate_start_time = time.time() try: retrying = AsyncRetrying( retry=retry_if_exception_type(EmptyModelOutputError), @@ -583,6 +587,16 @@ async def _iter_llm_responses_with_fallback( and (not is_last_candidate) ): last_err_response = resp + if resp.usage is not None: + self.stats.token_usage += resp.usage + self.provider_stat_segments.append( + ProviderStatSegment( + provider=candidate, + usage=resp.usage, + start_time=candidate_start_time, + end_time=time.time(), + ) + ) logger.warning( "Chat Model %s returns error response, trying fallback to next provider.", candidate_id, @@ -613,6 +627,18 @@ async def _iter_llm_responses_with_fallback( return except Exception as exc: # noqa: BLE001 last_exception = exc + failed_usage = getattr(exc, "_astrbot_token_usage", None) + if isinstance(failed_usage, TokenUsage): + self.stats.token_usage += failed_usage + if not is_last_candidate: + self.provider_stat_segments.append( + ProviderStatSegment( + provider=candidate, + usage=failed_usage, + start_time=candidate_start_time, + end_time=time.time(), + ) + ) logger.warning( "Chat Model %s request error: %s", candidate_id, diff --git a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py index b58e3a7abc..5053c41725 100644 --- a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py +++ b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py @@ -221,7 +221,9 @@ async def process( async with session_lock_manager.acquire_lock(event.unified_msg_origin): logger.debug("acquired session lock for llm request") agent_runner: AgentRunner | None = None + req: ProviderRequest | None = None runner_registered = False + stats_scheduled = False try: build_cfg = replace( self.main_agent_cfg, @@ -395,13 +397,12 @@ async def process( resp=final_resp.completion_text if final_resp else None, ) - asyncio.create_task( - _record_internal_agent_stats( - event, - req, - agent_runner, - final_resp, - ) + stats_scheduled = _schedule_internal_agent_stats( + stats_scheduled, + event, + req, + agent_runner, + final_resp, ) # 检查事件是否被停止,如果被停止则不保存历史记录 @@ -423,6 +424,14 @@ async def process( ), ) finally: + if agent_runner is not None: + stats_scheduled = _schedule_internal_agent_stats( + stats_scheduled, + event, + req, + agent_runner, + agent_runner.get_final_llm_resp(), + ) if runner_registered and agent_runner is not None: unregister_active_runner(event.unified_msg_origin, agent_runner) @@ -558,3 +567,18 @@ async def _record_internal_agent_stats( agent_runner=agent_runner, final_response=final_resp, ) + + +def _schedule_internal_agent_stats( + already_scheduled: bool, + event: AstrMessageEvent, + req: ProviderRequest | None, + agent_runner: AgentRunner | None, + final_resp: LLMResponse | None, +) -> bool: + if already_scheduled or agent_runner is None: + return already_scheduled + asyncio.create_task( + _record_internal_agent_stats(event, req, agent_runner, final_resp) + ) + return True diff --git a/astrbot/core/provider/stats.py b/astrbot/core/provider/stats.py index 26c61e19c1..872237e443 100644 --- a/astrbot/core/provider/stats.py +++ b/astrbot/core/provider/stats.py @@ -1,5 +1,6 @@ from __future__ import annotations +from dataclasses import dataclass from typing import Any from astrbot import logger @@ -7,6 +8,15 @@ from astrbot.core.provider.entities import LLMResponse, ProviderRequest, TokenUsage +@dataclass(slots=True) +class ProviderStatSegment: + provider: Any + usage: TokenUsage + start_time: float + end_time: float + status: str = "error" + + def _provider_id(provider: Any) -> str: provider_config = getattr(provider, "provider_config", {}) or {} return provider_config.get("id", "") or provider.meta().id @@ -50,6 +60,47 @@ async def record_agent_runner_stats( if request is not None and request.conversation is not None else None ) + segments: list[ProviderStatSegment] = list( + getattr(agent_runner, "provider_stat_segments", ()) + ) + segmented_usage = TokenUsage() + for segment in segments: + segmented_usage += segment.usage + await db.insert_provider_stat( + umo=umo, + conversation_id=conversation_id, + provider_id=_provider_id(segment.provider), + provider_model=segment.provider.get_model(), + status=segment.status, + stats={ + "token_usage": segment.usage.__dict__.copy(), + "start_time": segment.start_time, + "end_time": segment.end_time, + "time_to_first_token": 0.0, + }, + agent_type=agent_type, + ) + + aggregate_stats = stats.to_dict() + aggregate_usage = stats.token_usage - segmented_usage + aggregate_stats["token_usage"] = { + "input_other": max(0, aggregate_usage.input_other), + "input_cached": max(0, aggregate_usage.input_cached), + "output": max(0, aggregate_usage.output), + } + if segments: + original_start = aggregate_stats["start_time"] + aggregate_start = max( + original_start, + max(segment.end_time for segment in segments), + ) + aggregate_stats["start_time"] = aggregate_start + aggregate_stats["time_to_first_token"] = max( + 0.0, + aggregate_stats["time_to_first_token"] + - (aggregate_start - original_start), + ) + await db.insert_provider_stat( umo=umo, conversation_id=conversation_id, @@ -59,7 +110,7 @@ async def record_agent_runner_stats( final_response, agent_runner.was_aborted(), ), - stats=stats.to_dict(), + stats=aggregate_stats, agent_type=agent_type, ) except Exception as exc: # noqa: BLE001 From 505191619b57aea51b9a114914e31f7a563a833e Mon Sep 17 00:00:00 2001 From: Rhonin Date: Sat, 15 Aug 2026 15:16:20 +0800 Subject: [PATCH 7/8] fix: address provider stats review feedback --- .../agent/runners/tool_loop_agent_runner.py | 41 +++++----- astrbot/core/provider/stats.py | 12 ++- tests/test_tool_loop_agent_runner.py | 56 +++++++++++++ tests/unit/test_provider_stats.py | 78 ++++++++++++++++++- tests/unit/test_star_context.py | 56 +++++++++++++ 5 files changed, 221 insertions(+), 22 deletions(-) diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index 5b361408ff..3dd4888dfc 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -587,16 +587,17 @@ async def _iter_llm_responses_with_fallback( and (not is_last_candidate) ): last_err_response = resp - if resp.usage is not None: - self.stats.token_usage += resp.usage - self.provider_stat_segments.append( - ProviderStatSegment( - provider=candidate, - usage=resp.usage, - start_time=candidate_start_time, - end_time=time.time(), - ) + last_exception = None + failed_usage = resp.usage or TokenUsage() + self.stats.token_usage += failed_usage + self.provider_stat_segments.append( + ProviderStatSegment( + provider=candidate, + usage=failed_usage, + start_time=candidate_start_time, + end_time=time.time(), ) + ) logger.warning( "Chat Model %s returns error response, trying fallback to next provider.", candidate_id, @@ -627,18 +628,20 @@ async def _iter_llm_responses_with_fallback( return except Exception as exc: # noqa: BLE001 last_exception = exc + last_err_response = None failed_usage = getattr(exc, "_astrbot_token_usage", None) - if isinstance(failed_usage, TokenUsage): - self.stats.token_usage += failed_usage - if not is_last_candidate: - self.provider_stat_segments.append( - ProviderStatSegment( - provider=candidate, - usage=failed_usage, - start_time=candidate_start_time, - end_time=time.time(), - ) + if not isinstance(failed_usage, TokenUsage): + failed_usage = TokenUsage() + self.stats.token_usage += failed_usage + if not is_last_candidate: + self.provider_stat_segments.append( + ProviderStatSegment( + provider=candidate, + usage=failed_usage, + start_time=candidate_start_time, + end_time=time.time(), ) + ) logger.warning( "Chat Model %s request error: %s", candidate_id, diff --git a/astrbot/core/provider/stats.py b/astrbot/core/provider/stats.py index 872237e443..3ef8969441 100644 --- a/astrbot/core/provider/stats.py +++ b/astrbot/core/provider/stats.py @@ -36,6 +36,14 @@ def _runner_status(response: LLMResponse | None, aborted: bool) -> str: return "completed" +def _token_usage_dict(usage: TokenUsage) -> dict[str, int]: + return { + "input_other": usage.input_other, + "input_cached": usage.input_cached, + "output": usage.output, + } + + async def record_agent_runner_stats( db: BaseDatabase, *, @@ -73,7 +81,7 @@ async def record_agent_runner_stats( provider_model=segment.provider.get_model(), status=segment.status, stats={ - "token_usage": segment.usage.__dict__.copy(), + "token_usage": _token_usage_dict(segment.usage), "start_time": segment.start_time, "end_time": segment.end_time, "time_to_first_token": 0.0, @@ -138,7 +146,7 @@ async def record_llm_response_stats( provider_model=provider.get_model(), status=_response_status(response), stats={ - "token_usage": usage.__dict__.copy(), + "token_usage": _token_usage_dict(usage), "start_time": start_time, "end_time": end_time, "time_to_first_token": 0.0, diff --git a/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index e316bd9b12..550d246430 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -195,6 +195,16 @@ async def text_chat(self, **kwargs) -> LLMResponse: raise error +class MockUsageErrProvider(MockProvider): + async def text_chat(self, **kwargs) -> LLMResponse: + self.call_count += 1 + return LLMResponse( + role="err", + completion_text="provider returned error", + usage=TokenUsage(input_other=3, input_cached=2, output=1), + ) + + class MockErrProvider(MockProvider): async def text_chat(self, **kwargs) -> LLMResponse: self.call_count += 1 @@ -1223,6 +1233,10 @@ async def test_fallback_provider_used_when_primary_raises( assert final_resp.completion_text == "这是我的最终回答" assert primary_provider.call_count == 1 assert fallback_provider.call_count == 1 + assert len(runner.provider_stat_segments) == 1 + segment = runner.provider_stat_segments[0] + assert segment.provider is primary_provider + assert segment.usage == TokenUsage() @pytest.mark.asyncio @@ -1289,6 +1303,48 @@ async def test_fallback_provider_used_when_primary_returns_err( assert final_resp.completion_text == "这是我的最终回答" assert primary_provider.call_count == 1 assert fallback_provider.call_count == 1 + assert len(runner.provider_stat_segments) == 1 + segment = runner.provider_stat_segments[0] + assert segment.provider is primary_provider + assert segment.usage == TokenUsage() + + +@pytest.mark.asyncio +async def test_fallback_consecutive_failures_do_not_duplicate_prior_usage( + runner, + provider_request, + mock_tool_executor, + mock_hooks, +): + primary_provider = MockUsageErrProvider() + fallback_provider = MockUsageFailingProvider() + + await runner.reset( + provider=primary_provider, + request=provider_request, + run_context=ContextWrapper(context=None), + tool_executor=mock_tool_executor, + agent_hooks=mock_hooks, + streaming=False, + fallback_providers=[fallback_provider], + ) + + async for _ in runner.step_until_done(5): + pass + + final_resp = runner.get_final_llm_resp() + assert final_resp is not None + assert final_resp.role == "err" + assert "RuntimeError" in final_resp.completion_text + assert runner.stats.token_usage == TokenUsage( + input_other=11, + input_cached=6, + output=7, + ) + assert len(runner.provider_stat_segments) == 1 + segment = runner.provider_stat_segments[0] + assert segment.provider is primary_provider + assert segment.usage == TokenUsage(input_other=3, input_cached=2, output=1) @pytest.mark.asyncio diff --git a/tests/unit/test_provider_stats.py b/tests/unit/test_provider_stats.py index cb0a4c4b45..746567fd1c 100644 --- a/tests/unit/test_provider_stats.py +++ b/tests/unit/test_provider_stats.py @@ -9,7 +9,83 @@ from astrbot.core.db.po import ProviderStat from astrbot.core.pipeline.process_stage.method.agent_sub_stages import internal from astrbot.core.provider.entities import ProviderRequest, TokenUsage -from astrbot.core.provider.stats import ProviderStatSegment +from astrbot.core.provider.stats import ( + ProviderStatSegment, + record_agent_runner_stats, + record_llm_response_stats, +) + + +def assert_public_token_usage(token_usage: dict[str, int]) -> None: + assert token_usage == { + "input_other": 5, + "input_cached": 2, + "output": 3, + } + + +@pytest.mark.asyncio +async def test_record_llm_response_stats_only_passes_public_token_fields(): + usage = TokenUsage(input_other=5, input_cached=2, output=3) + usage.internal_note = "must not reach the database" + db = SimpleNamespace(insert_provider_stat=AsyncMock()) + provider = SimpleNamespace( + provider_config={"id": "provider-1"}, + meta=lambda: SimpleNamespace(id="provider-1"), + get_model=lambda: "test-model", + ) + + await record_llm_response_stats( + db, + umo="provider:provider-1:test", + provider=provider, + response=SimpleNamespace(role="assistant", usage=usage), + start_time=100.0, + end_time=101.0, + ) + + stats = db.insert_provider_stat.await_args.kwargs["stats"] + assert_public_token_usage(stats["token_usage"]) + + +@pytest.mark.asyncio +async def test_record_agent_runner_stats_only_passes_public_segment_token_fields(): + usage = TokenUsage(input_other=5, input_cached=2, output=3) + usage.internal_note = "must not reach the database" + db = SimpleNamespace(insert_provider_stat=AsyncMock()) + provider = SimpleNamespace( + provider_config={"id": "provider-1"}, + meta=lambda: SimpleNamespace(id="provider-1"), + get_model=lambda: "test-model", + ) + runner = SimpleNamespace( + provider=provider, + stats=AgentStats( + token_usage=usage, + start_time=100.0, + end_time=102.0, + ), + provider_stat_segments=[ + ProviderStatSegment( + provider=provider, + usage=usage, + start_time=100.0, + end_time=101.0, + ) + ], + was_aborted=lambda: False, + ) + + await record_agent_runner_stats( + db, + umo="test:session", + request=None, + agent_runner=runner, + final_response=SimpleNamespace(role="assistant"), + ) + + segment_stats = db.insert_provider_stat.await_args_list[0].kwargs["stats"] + assert_public_token_usage(segment_stats["token_usage"]) @pytest.mark.asyncio diff --git a/tests/unit/test_star_context.py b/tests/unit/test_star_context.py index 9a7862d3a2..4eb72bd534 100644 --- a/tests/unit/test_star_context.py +++ b/tests/unit/test_star_context.py @@ -177,6 +177,62 @@ async def test_llm_generate_persists_one_provider_stat(temp_db): assert record.end_time >= record.start_time > 0 +@pytest.mark.asyncio +async def test_llm_generate_persists_error_stat_when_provider_raises(temp_db): + provider = StatsProvider() + provider.text_chat = AsyncMock(side_effect=RuntimeError("provider failed")) + context = Context.__new__(Context) + context._db = temp_db + context.provider_manager = SimpleNamespace( + get_provider_by_id=AsyncMock(return_value=provider), + ) + + with pytest.raises(RuntimeError, match="provider failed"): + await context.llm_generate( + chat_provider_id="provider-1", + prompt="test", + session_id="failed-session", + ) + + records = await get_provider_stats(temp_db) + assert len(records) == 1 + record = records[0] + assert record.status == "error" + assert record.token_input_other == 0 + assert record.token_input_cached == 0 + assert record.token_output == 0 + + +@pytest.mark.asyncio +async def test_llm_generate_persists_error_response_usage(temp_db): + provider = StatsProvider() + error_response = LLMResponse( + role="err", + usage=TokenUsage(input_other=7, input_cached=2, output=1), + ) + provider.text_chat = AsyncMock(return_value=error_response) + context = Context.__new__(Context) + context._db = temp_db + context.provider_manager = SimpleNamespace( + get_provider_by_id=AsyncMock(return_value=provider), + ) + + response = await context.llm_generate( + chat_provider_id="provider-1", + prompt="test", + session_id="error-response-session", + ) + + records = await get_provider_stats(temp_db) + assert response is error_response + assert len(records) == 1 + record = records[0] + assert record.status == "error" + assert record.token_input_other == 7 + assert record.token_input_cached == 2 + assert record.token_output == 1 + + @pytest.mark.asyncio async def test_tool_loop_agent_persists_one_aggregated_provider_stat( temp_db, From 1a3e0dbb44c6a3c10b623687bbece92ffd40243f Mon Sep 17 00:00:00 2001 From: Rhonin Date: Sat, 22 Aug 2026 23:59:07 +0800 Subject: [PATCH 8/8] fix: support anthropic SDK 1.0 httpx2 module (#9769) --- astrbot/core/provider/sources/anthropic_source.py | 6 +++++- tests/test_anthropic_kimi_code_provider.py | 8 +++++++- 2 files changed, 12 insertions(+), 2 deletions(-) diff --git a/astrbot/core/provider/sources/anthropic_source.py b/astrbot/core/provider/sources/anthropic_source.py index 27cc459622..e9dfbc3d8d 100644 --- a/astrbot/core/provider/sources/anthropic_source.py +++ b/astrbot/core/provider/sources/anthropic_source.py @@ -131,7 +131,11 @@ def _create_http_client(self, provider_config: dict) -> httpx.AsyncClient | None try: from anthropic import _base_client as anthropic_base_client - httpx_module = getattr(anthropic_base_client, "httpx", httpx) + httpx_module = getattr( + anthropic_base_client, + "httpx2", + getattr(anthropic_base_client, "httpx", httpx), + ) except ImportError: pass return create_proxy_client( diff --git a/tests/test_anthropic_kimi_code_provider.py b/tests/test_anthropic_kimi_code_provider.py index 0dc33f58ba..d54894a05f 100644 --- a/tests/test_anthropic_kimi_code_provider.py +++ b/tests/test_anthropic_kimi_code_provider.py @@ -137,7 +137,13 @@ def fake_create_proxy_client( assert captured["provider_label"] == "Anthropic" assert captured["proxy"] == "http://127.0.0.1:7890" assert captured["headers"] == {"X-Trace-Id": "trace-1"} - assert captured["httpx_module"] is anthropic_base_client.httpx + sdk_httpx = getattr( + anthropic_base_client, + "httpx2", + getattr(anthropic_base_client, "httpx", None), + ) + assert sdk_httpx is not None + assert captured["httpx_module"] is sdk_httpx def test_create_http_client_falls_back_to_global_httpx_module(monkeypatch):