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
122 changes: 112 additions & 10 deletions astrbot/core/agent/runners/tool_loop_agent_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
import traceback
import typing as T
import uuid
from contextlib import suppress
from contextlib import contextmanager, suppress
from dataclasses import dataclass, field, replace
from pathlib import Path

Expand Down 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.provider import Provider, provider_stats_managed_by_agent
from astrbot.core.provider.stats import ProviderStatSegment

from ..context.compressor import ContextCompressor
from ..context.config import ContextConfig
Expand Down Expand Up @@ -227,6 +229,7 @@ async def reset(
tool_schema_mode: str | None = "full",
fallback_providers: list[Provider] | None = None,
request_max_retries: int | None = None,
provider_stats_managed_by_agent: bool = False,
tool_result_overflow_dir: str | None = None,
read_tool: FunctionTool | None = None,
**kwargs: T.Any,
Expand All @@ -241,6 +244,7 @@ async def reset(
self.custom_token_counter = custom_token_counter
self.custom_compressor = custom_compressor
self.request_max_retries = request_max_retries
self.provider_stats_managed_by_agent = provider_stats_managed_by_agent
self.tool_result_overflow_dir = tool_result_overflow_dir
self.read_tool = read_tool
self._tool_result_token_counter = EstimateTokenCounter()
Expand Down Expand Up @@ -324,6 +328,9 @@ async def reset(

self.stats = AgentStats()
self.stats.start_time = time.time()
self.provider_stat_segments: list[ProviderStatSegment] = []
self._provider_token_usage: dict[int, TokenUsage] = {}
self._provider_usage_start_times: dict[int, float] = {}

def _read_tool_hint(self) -> str:
if self.read_tool is not None:
Expand Down Expand Up @@ -458,6 +465,76 @@ def _truncate_tool_result_preview(
preview = preview[:next_len]
return preview

@contextmanager
def _provider_stats_scope(self) -> T.Iterator[None]:
token = provider_stats_managed_by_agent.set(
self.provider_stats_managed_by_agent
)
try:
yield
finally:
provider_stats_managed_by_agent.reset(token)

def _mark_provider_attempt_started(
self,
provider: Provider,
start_time: float,
) -> float:
provider_key = id(provider)
return self._provider_usage_start_times.setdefault(provider_key, start_time)

def _accumulate_token_usage(
self,
usage: TokenUsage | None,
provider: Provider | None = None,
) -> None:
if usage is None:
return
self.stats.token_usage += usage
usage_provider = provider or self.provider
provider_key = id(usage_provider)
provider_usage = self._provider_token_usage.get(provider_key, TokenUsage())
self._provider_token_usage[provider_key] = provider_usage + usage
self.stats.current_context_tokens = usage.input
if self.req and self.req.conversation:
self.req.conversation.token_usage = usage.total

def _settle_provider_stat_segment(
self,
provider: Provider,
*,
fallback_start_time: float,
end_time: float,
) -> None:
provider_key = id(provider)
self.provider_stat_segments.append(
ProviderStatSegment(
provider=provider,
usage=self._provider_token_usage.pop(provider_key, TokenUsage()),
start_time=self._provider_usage_start_times.pop(
provider_key,
fallback_start_time,
),
end_time=end_time,
)
)

async def _await_additional_provider_response(
self,
awaitable: T.Awaitable[LLMResponse],
) -> LLMResponse | None:
with self._provider_stats_scope():
try:
response = await self._await_or_stop(awaitable)
except Exception as exc:
failed_usage = getattr(exc, "_astrbot_token_usage", None)
if isinstance(failed_usage, TokenUsage):
self._accumulate_token_usage(failed_usage)
raise
if response is not None:
self._accumulate_token_usage(response.usage)
return response

async def _await_or_stop(
self,
awaitable: T.Awaitable[AwaitableResultT],
Expand Down Expand Up @@ -517,7 +594,8 @@ async def _iter_llm_responses(
try:
while True:
try:
resp = await self._await_or_stop(anext(stream)) # type: ignore
with self._provider_stats_scope():
resp = await self._await_or_stop(anext(stream)) # type: ignore
except StopAsyncIteration:
return
if resp is None:
Expand All @@ -526,7 +604,8 @@ async def _iter_llm_responses(
finally:
await self._close_executor(stream)
else:
resp = await self._await_or_stop(self.provider.text_chat(**payload))
with self._provider_stats_scope():
resp = await self._await_or_stop(self.provider.text_chat(**payload))
if resp is not None:
yield resp

Expand Down Expand Up @@ -562,6 +641,10 @@ async def _iter_llm_responses_with_fallback(
candidate_id,
)
self.provider = candidate
candidate_start_time = self._mark_provider_attempt_started(
candidate,
time.time(),
)
try:
retrying = AsyncRetrying(
retry=retry_if_exception_type(EmptyModelOutputError),
Expand Down Expand Up @@ -594,6 +677,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._accumulate_token_usage(
failed_usage,
candidate,
)
self._settle_provider_stat_segment(
candidate,
fallback_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 @@ -624,6 +718,17 @@ 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._accumulate_token_usage(failed_usage, candidate)
if not is_last_candidate:
self._settle_provider_stat_segment(
candidate,
fallback_start_time=candidate_start_time,
end_time=time.time(),
)
logger.warning(
"Chat Model %s request error: %s",
candidate_id,
Expand Down Expand Up @@ -864,10 +969,7 @@ async def step(self):
if llm_response.usage:
# Keep cumulative usage for billing and expose the latest request
# input separately for context-window occupancy displays.
self.stats.token_usage += llm_response.usage
self.stats.current_context_tokens = llm_response.usage.input
if self.req.conversation:
self.req.conversation.token_usage = llm_response.usage.total
self._accumulate_token_usage(llm_response.usage)
# end_time must be set before the yield serializes to_dict().
self.stats.end_time = time.time()
yield AgentResponse(
Expand Down Expand Up @@ -1441,7 +1543,7 @@ async def _resolve_tool_exec(
)
if param_subset.tools and tool_names:
contexts = self._build_tool_requery_context(tool_names)
requery_resp = await self._await_or_stop(
requery_resp = await self._await_additional_provider_response(
self.provider.text_chat(
contexts=self._sanitize_contexts_for_provider(contexts),
func_tool=param_subset,
Expand Down Expand Up @@ -1471,7 +1573,7 @@ async def _resolve_tool_exec(
tool_names,
extra_instruction=self.SKILLS_LIKE_REQUERY_REPAIR_INSTRUCTION,
)
repair_resp = await self._await_or_stop(
repair_resp = await self._await_additional_provider_response(
self.provider.text_chat(
contexts=self._sanitize_contexts_for_provider(
repair_contexts
Expand Down
1 change: 1 addition & 0 deletions astrbot/core/astr_main_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -1898,6 +1898,7 @@ async def build_main_agent(
tool_schema_mode=config.tool_schema_mode,
fallback_providers=fallback_providers,
request_max_retries=config.request_max_retries,
provider_stats_managed_by_agent=True,
tool_result_overflow_dir=(
get_astrbot_system_tmp_path()
if req.func_tool and req.func_tool.get_tool("astrbot_file_read_tool")
Expand Down
35 changes: 21 additions & 14 deletions astrbot/core/cron/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,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.config_number import coerce_int_config
from astrbot.core.utils.history_saver import persist_agent_history

Expand Down Expand Up @@ -508,21 +509,27 @@ async def _woke_main_agent(
raise RuntimeError("Failed to build main agent for cron job.")

runner = result.agent_runner
async for _ in runner.step_until_done(agent_max_step):
# agent will send message to user via using tools
pass
llm_resp = runner.get_final_llm_resp()
if runner.state == AgentState.ERROR:
# The run failed (e.g. malformed function call at max steps) but
# no exception escapes the runner; without this the job was
# recorded as completed with last_error=NULL and the user saw
# only intermediate messages (#9980).
detail = (
f": {llm_resp.completion_text}"
if llm_resp and llm_resp.completion_text
else ""
llm_resp = None
try:
async for _ in runner.step_until_done(agent_max_step):
# agent will send message to user via using tools
pass
llm_resp = runner.get_final_llm_resp()
if getattr(runner, "state", None) == AgentState.ERROR:
detail = (
f": {llm_resp.completion_text}"
if llm_resp and llm_resp.completion_text
else ""
)
raise RuntimeError(f"Cron agent run ended in ERROR state{detail}")
finally:
await record_agent_runner_stats(
self.db,
umo=cron_event.unified_msg_origin,
request=req,
agent_runner=runner,
final_response=llm_resp,
)
raise RuntimeError(f"Cron agent run ended in ERROR state{detail}")
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
Loading
Loading