Skip to content
29 changes: 29 additions & 0 deletions astrbot/core/agent/runners/tool_loop_agent_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,13 +39,15 @@
from astrbot.core.provider.entities import (
LLMResponse,
ProviderRequest,
TokenUsage,
ToolCallsResult,
)
from astrbot.core.provider.modalities import (
log_context_sanitize_stats,
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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
19 changes: 15 additions & 4 deletions astrbot/core/cron/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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', '')} "
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
)

# 检查事件是否被停止,如果被停止则不保存历史记录
Expand All @@ -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)

Expand Down Expand Up @@ -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
6 changes: 5 additions & 1 deletion astrbot/core/provider/sources/anthropic_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
157 changes: 157 additions & 0 deletions astrbot/core/provider/stats.py
Original file line number Diff line number Diff line change
@@ -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,
Comment thread
sourcery-ai[bot] marked this conversation as resolved.
"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)
Loading
Loading