diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index 8c91adbbfd..3dd4888dfc 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,17 @@ async def _iter_llm_responses_with_fallback( and (not is_last_candidate) ): last_err_response = resp + 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, @@ -613,6 +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 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/cron/manager.py b/astrbot/core/cron/manager.py index b5a0e7c3e4..0fc9543891 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: @@ -488,10 +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() + 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/pipeline/process_stage/method/agent_sub_stages/internal.py b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py index 40e0e99a50..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 @@ -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 @@ -220,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, @@ -394,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, ) # 检查事件是否被停止,如果被停止则不保存历史记录 @@ -422,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) @@ -550,37 +560,25 @@ 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 - ) + await record_agent_runner_stats( + db_helper, + umo=event.unified_msg_origin, + request=req, + agent_runner=agent_runner, + final_response=final_resp, + ) - 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) + +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/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/astrbot/core/provider/stats.py b/astrbot/core/provider/stats.py new file mode 100644 index 0000000000..3ef8969441 --- /dev/null +++ b/astrbot/core/provider/stats.py @@ -0,0 +1,157 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +from astrbot import logger +from astrbot.core.db import BaseDatabase +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 + + +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 None or response.role == "err": + return "error" + 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, + *, + 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 + ) + 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": _token_usage_dict(segment.usage), + "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, + provider_id=_provider_id(provider), + provider_model=provider.get_model(), + status=_runner_status( + final_response, + agent_runner.was_aborted(), + ), + stats=aggregate_stats, + 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": _token_usage_dict(usage), + "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..f4b9679b7b 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,30 @@ 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:{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 async def tool_loop_agent( @@ -319,9 +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() + 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 3a9f5f80d5..be642e07d0 100644 --- a/astrbot/dashboard/services/stat_service.py +++ b/astrbot/dashboard/services/stat_service.py @@ -320,7 +320,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()) @@ -374,7 +374,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 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): diff --git a/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index 1e679de4aa..550d246430 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -183,6 +183,28 @@ 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 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 @@ -1211,6 +1233,47 @@ 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 +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 @@ -1240,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_cron_manager.py b/tests/unit/test_cron_manager.py index 47f97ef445..7b31976df9 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,167 @@ 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 + + @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_provider_stats.py b/tests/unit/test_provider_stats.py index c306c3f820..746567fd1c 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,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, + 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 @@ -63,3 +142,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) diff --git a/tests/unit/test_star_context.py b/tests/unit/test_star_context.py index 0979aaa0bf..4eb72bd534 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,211 @@ 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 == "provider" + 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_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, + 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 + + +@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 index b024f3a9b4..c9a57a7cc8 100644 --- a/tests/unit/test_stat_service.py +++ b/tests/unit/test_stat_service.py @@ -1,5 +1,6 @@ import time from datetime import datetime, timedelta +from types import SimpleNamespace from unittest.mock import MagicMock import pytest @@ -78,3 +79,32 @@ async def test_get_stat_empty_window(temp_db): assert result["platform"] == [] assert result["message_count"] == 4 assert all(count == 0 for _, count in result["message_time_series"]) + + +@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)