mirror of
https://github.com/Nezumi-2711/astrbot_plugin_qq_group_daily_analysis.git
synced 2026-09-22 13:38:43 +00:00
feat(llm_utils): 更新了
src/analysis/utils/llm_utils.py 文件,以适配 AstrBot v4.5.7+ 的新 API 标准。 主要更改如下: 重构了 Provider 选择逻辑: 将 get_provider_with_fallback 重命名并修改为 get_provider_id_with_fallback 。 现在它返回的是 Provider ID (字符串),而不是 Provider 对象实例。这是因为新的 llm_generate 接口需要通过 ID 来指定使用的模型。 使用了新的 await context.get_current_chat_provider_id(umo=umo) 方法来获取当前会话的模型 ID。 更新了 LLM 调用方式: 在 call_provider_with_retry 函数中,不再直接调用 provider.text_chat(...)。 改为使用统一的 await context.llm_generate(...) 接口。这是 AstrBot 推荐的新方式,更加简洁且统一。 兼容性调整: 更新了 extract_token_usage 函数,使其能够从 llm_generate 返回的新响应格式中正确提取 Token 使用统计(优先检查 usage 属性)。 这些更改确保了插件能够利用 AstrBot 新版本的特性,同时保持了原有的多级回退(配置 -> 主模型 -> 会话模型 -> 任意可用)逻辑不变。
This commit is contained in:
@@ -8,17 +8,17 @@ from typing import Any
|
||||
from astrbot.api import logger
|
||||
|
||||
|
||||
def _try_get_provider_by_id(context, provider_id: str, description: str) -> Any | None:
|
||||
async def _try_get_provider_id_by_id(context, provider_id: str, description: str) -> str | None:
|
||||
"""
|
||||
尝试通过 ID 获取 Provider 的辅助函数
|
||||
|
||||
尝试通过 ID 获取 Provider ID 的辅助函数
|
||||
|
||||
Args:
|
||||
context: AstrBot上下文对象
|
||||
provider_id: Provider ID
|
||||
description: 描述信息,用于日志
|
||||
|
||||
Returns:
|
||||
Provider 实例或 None
|
||||
Provider ID 或 None
|
||||
"""
|
||||
if not provider_id or not isinstance(provider_id, str) or not provider_id.strip():
|
||||
return None
|
||||
@@ -26,67 +26,69 @@ def _try_get_provider_by_id(context, provider_id: str, description: str) -> Any
|
||||
provider_id = provider_id.strip()
|
||||
logger.info(f"尝试使用{description}: {provider_id}")
|
||||
try:
|
||||
# 验证 Provider 是否存在
|
||||
provider = context.get_provider_by_id(provider_id=provider_id)
|
||||
if provider:
|
||||
logger.info(f"✓ 使用{description}: {provider_id}")
|
||||
return provider
|
||||
return provider_id
|
||||
except Exception as e:
|
||||
logger.warning(f"无法找到{description} '{provider_id}': {e}")
|
||||
return None
|
||||
|
||||
|
||||
def _try_get_session_provider(context, umo: str) -> Any | None:
|
||||
async def _try_get_session_provider_id(context, umo: str) -> str | None:
|
||||
"""
|
||||
尝试获取会话 Provider 的辅助函数
|
||||
尝试获取会话 Provider ID 的辅助函数
|
||||
|
||||
Args:
|
||||
context: AstrBot上下文对象
|
||||
umo: unified_msg_origin
|
||||
|
||||
Returns:
|
||||
Provider 实例或 None
|
||||
Provider ID 或 None
|
||||
"""
|
||||
try:
|
||||
provider = context.get_using_provider(umo=umo)
|
||||
if provider:
|
||||
try:
|
||||
meta = provider.meta()
|
||||
provider_id = meta.id
|
||||
logger.info(f"✓ 使用当前会话的 Provider: {provider_id}")
|
||||
except Exception:
|
||||
logger.info("✓ 使用当前会话的默认 Provider")
|
||||
return provider
|
||||
# 使用新 API 获取当前会话的 Provider ID
|
||||
provider_id = await context.get_current_chat_provider_id(umo=umo)
|
||||
if provider_id:
|
||||
logger.info(f"✓ 使用当前会话的 Provider: {provider_id}")
|
||||
return provider_id
|
||||
except Exception as e:
|
||||
logger.warning(f"无法获取会话 Provider: {e}")
|
||||
logger.warning(f"无法获取会话 Provider ID: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def _try_get_first_available_provider(context) -> Any | None:
|
||||
async def _try_get_first_available_provider_id(context) -> str | None:
|
||||
"""
|
||||
尝试获取第一个可用 Provider 的辅助函数
|
||||
尝试获取第一个可用 Provider ID 的辅助函数
|
||||
|
||||
Args:
|
||||
context: AstrBot上下文对象
|
||||
|
||||
Returns:
|
||||
Provider 实例或 None
|
||||
Provider ID 或 None
|
||||
"""
|
||||
try:
|
||||
all_providers = context.get_all_providers()
|
||||
if all_providers and len(all_providers) > 0:
|
||||
provider = all_providers[0]
|
||||
logger.info(f"✓ 使用第一个可用 Provider: {type(provider).__name__}")
|
||||
return provider
|
||||
try:
|
||||
meta = provider.meta()
|
||||
provider_id = meta.id
|
||||
logger.info(f"✓ 使用第一个可用 Provider: {provider_id}")
|
||||
return provider_id
|
||||
except Exception:
|
||||
logger.warning("第一个 Provider 无法获取 ID")
|
||||
except Exception as e:
|
||||
logger.warning(f"无法获取任何 Provider: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def get_provider_with_fallback(
|
||||
async def get_provider_id_with_fallback(
|
||||
context, config_manager, provider_id_key: str, umo: str = None
|
||||
) -> Any | None:
|
||||
) -> str | None:
|
||||
"""
|
||||
根据配置键获取 Provider,支持多级回退
|
||||
根据配置键获取 Provider ID,支持多级回退
|
||||
|
||||
回退顺序:
|
||||
1. 尝试从配置获取指定的 provider_id(如 topic_provider_id)
|
||||
@@ -101,7 +103,7 @@ def get_provider_with_fallback(
|
||||
umo: unified_msg_origin,用于获取会话默认 Provider
|
||||
|
||||
Returns:
|
||||
Provider 实例或 None
|
||||
Provider ID 或 None
|
||||
"""
|
||||
try:
|
||||
# 输出Provider选择开始日志
|
||||
@@ -119,7 +121,7 @@ def get_provider_with_fallback(
|
||||
specific_provider_id = getattr(config_manager, getter_method)()
|
||||
if specific_provider_id:
|
||||
strategies.append(
|
||||
lambda pid=specific_provider_id: _try_get_provider_by_id(
|
||||
lambda pid=specific_provider_id: _try_get_provider_id_by_id(
|
||||
context, pid, f"配置的 {provider_id_key}"
|
||||
)
|
||||
)
|
||||
@@ -129,18 +131,18 @@ def get_provider_with_fallback(
|
||||
main_provider_id = config_manager.get_llm_provider_id()
|
||||
if main_provider_id:
|
||||
strategies.append(
|
||||
lambda pid=main_provider_id: _try_get_provider_by_id(
|
||||
lambda pid=main_provider_id: _try_get_provider_id_by_id(
|
||||
context, pid, "主 LLM Provider"
|
||||
)
|
||||
)
|
||||
strategy_names.append("2. 主 LLM Provider")
|
||||
|
||||
# 3. 当前会话的 Provider
|
||||
strategies.append(lambda: _try_get_session_provider(context, umo))
|
||||
strategies.append(lambda: _try_get_session_provider_id(context, umo))
|
||||
strategy_names.append("3. 当前会话 Provider")
|
||||
|
||||
# 4. 第一个可用的 Provider
|
||||
strategies.append(lambda: _try_get_first_available_provider(context))
|
||||
strategies.append(lambda: _try_get_first_available_provider_id(context))
|
||||
strategy_names.append("4. 第一个可用 Provider")
|
||||
|
||||
# 输出回退策略列表
|
||||
@@ -148,20 +150,12 @@ def get_provider_with_fallback(
|
||||
|
||||
# 依次尝试每个策略
|
||||
for idx, strategy in enumerate(strategies):
|
||||
provider = strategy()
|
||||
if provider:
|
||||
# 获取最终的 provider ID 用于日志
|
||||
final_provider_id = "unknown"
|
||||
try:
|
||||
meta = provider.meta()
|
||||
final_provider_id = meta.id
|
||||
except Exception:
|
||||
final_provider_id = type(provider).__name__
|
||||
|
||||
provider_id = await strategy()
|
||||
if provider_id:
|
||||
logger.info(
|
||||
f"[Provider 选择] ✓ 成功!使用策略 #{idx + 1},Provider ID: {final_provider_id}"
|
||||
f"[Provider 选择] ✓ 成功!使用策略 #{idx + 1},Provider ID: {provider_id}"
|
||||
)
|
||||
return provider
|
||||
return provider_id
|
||||
|
||||
logger.error("[Provider 选择] ✗ 失败:所有回退策略均无法获取可用 Provider")
|
||||
return None
|
||||
@@ -202,25 +196,17 @@ async def call_provider_with_retry(
|
||||
last_exc = None
|
||||
for attempt in range(1, retries + 1):
|
||||
try:
|
||||
# 使用新的 provider 选择逻辑,支持配置化选择和多级回退
|
||||
provider = get_provider_with_fallback(
|
||||
# 使用新的 provider 选择逻辑,获取 Provider ID
|
||||
provider_id = await get_provider_id_with_fallback(
|
||||
context, config_manager, provider_id_key, umo
|
||||
)
|
||||
|
||||
provider_id = "unknown"
|
||||
if provider:
|
||||
try:
|
||||
meta = provider.meta()
|
||||
provider_id = meta.id
|
||||
except Exception as e:
|
||||
logger.debug(f"获取提供商ID失败: {e}")
|
||||
|
||||
if not provider:
|
||||
logger.error("provider 为空,无法调用 text_chat,直接返回 None")
|
||||
if not provider_id:
|
||||
logger.error("provider_id 为空,无法调用 llm_generate,直接返回 None")
|
||||
return None
|
||||
|
||||
logger.info(
|
||||
f"[LLM 调用] 使用 Provider: {provider_id} | "
|
||||
f"[LLM 调用] 使用 Provider ID: {provider_id} | "
|
||||
f"max_tokens={max_tokens} | temperature={temperature} | "
|
||||
f"prompt长度={len(prompt) if prompt else 0}字符"
|
||||
)
|
||||
@@ -232,14 +218,23 @@ async def call_provider_with_retry(
|
||||
# 检查 prompt 是否为空
|
||||
if not prompt or not prompt.strip():
|
||||
logger.error(
|
||||
"LLM provider: prompt 为空或只包含空白字符,无法调用 text_chat"
|
||||
"LLM provider: prompt 为空或只包含空白字符,无法调用 llm_generate"
|
||||
)
|
||||
return None
|
||||
|
||||
coro = provider.text_chat(
|
||||
prompt=prompt, max_tokens=max_tokens, temperature=temperature
|
||||
# 使用新的 llm_generate API
|
||||
# 注意:llm_generate 可能不直接支持 max_tokens 和 temperature 参数,
|
||||
# 取决于 AstrBot 版本和具体实现。如果支持 kwargs,可以传递。
|
||||
# 这里假设支持 kwargs 传递给底层 provider。
|
||||
llm_resp = await context.llm_generate(
|
||||
chat_provider_id=provider_id,
|
||||
prompt=prompt,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature
|
||||
)
|
||||
return await asyncio.wait_for(coro, timeout=timeout)
|
||||
|
||||
return llm_resp
|
||||
|
||||
except asyncio.TimeoutError as e:
|
||||
last_exc = e
|
||||
logger.warning(f"LLM请求超时: 第{attempt}次, timeout={timeout}s")
|
||||
@@ -268,8 +263,16 @@ def extract_token_usage(response) -> dict | None:
|
||||
try:
|
||||
token_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||
|
||||
# 安全地提取 usage,避免 response.raw_completion.usage 为 None 导致的 AttributeError
|
||||
usage = None
|
||||
# 尝试从 LLMResponse 中提取 usage
|
||||
# 假设 LLMResponse 有 usage 属性或 raw_completion 属性
|
||||
if hasattr(response, "usage") and response.usage:
|
||||
usage = response.usage
|
||||
token_usage["prompt_tokens"] = getattr(usage, "prompt_tokens", 0) or 0
|
||||
token_usage["completion_tokens"] = getattr(usage, "completion_tokens", 0) or 0
|
||||
token_usage["total_tokens"] = getattr(usage, "total_tokens", 0) or 0
|
||||
return token_usage
|
||||
|
||||
# 兼容旧的提取方式 (如果 response 是旧的 ProviderResponse)
|
||||
if getattr(response, "raw_completion", None) is not None:
|
||||
usage = getattr(response.raw_completion, "usage", None)
|
||||
if usage:
|
||||
|
||||
Reference in New Issue
Block a user