Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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
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