mirror of
https://github.com/Nezumi-2711/astrbot_plugin_qq_group_daily_analysis.git
synced 2026-09-22 20:01:04 +00:00
feat: 支持可选的流式 LLM 调用 (#181 @anchorAnc)
* feat: add optional streaming LLM provider call * fix(llm_utils): code quality --------- Co-authored-by: SXP-Simon <sxp20061207@163.com>
This commit is contained in:
co-authored by
SXP-Simon
parent
fd51b4e399
commit
83486de1d1
@@ -5,6 +5,8 @@ LLM API请求处理工具模块
|
||||
|
||||
import asyncio
|
||||
|
||||
from astrbot.api.provider import LLMResponse
|
||||
|
||||
from ....utils.logger import logger
|
||||
from ....utils.resilience import CircuitBreaker, GlobalRateLimiter
|
||||
from .structured_output_schema import JSONObject
|
||||
@@ -36,6 +38,40 @@ def _get_circuit_breaker(provider_id: str) -> CircuitBreaker:
|
||||
return _circuit_breakers[provider_id]
|
||||
|
||||
|
||||
async def _call_provider_stream(
|
||||
context, provider_id: str, llm_kwargs: dict[str, object]
|
||||
):
|
||||
provider = context.get_provider_by_id(provider_id=provider_id)
|
||||
if provider is None:
|
||||
raise RuntimeError(f"Provider 不存在: {provider_id}")
|
||||
|
||||
stream_kwargs = dict(llm_kwargs)
|
||||
stream_kwargs.pop("chat_provider_id", None)
|
||||
|
||||
final_resp = None
|
||||
content_parts: list[str] = []
|
||||
async for resp in provider.text_chat_stream(**stream_kwargs):
|
||||
final_resp = resp
|
||||
if getattr(resp, "is_chunk", False):
|
||||
text = getattr(resp, "completion_text", "")
|
||||
if text:
|
||||
content_parts.append(text)
|
||||
|
||||
if final_resp is None:
|
||||
raise RuntimeError("流式 LLM 调用未返回任何响应")
|
||||
|
||||
final_text = extract_response_text(final_resp)
|
||||
if final_text and not getattr(final_resp, "is_chunk", False):
|
||||
return final_resp
|
||||
|
||||
return LLMResponse(
|
||||
role="assistant",
|
||||
completion_text="".join(content_parts),
|
||||
usage=getattr(final_resp, "usage", None),
|
||||
raw_completion=getattr(final_resp, "raw_completion", None),
|
||||
)
|
||||
|
||||
|
||||
async def _try_get_provider_id_by_id(
|
||||
context, provider_id: str, description: str
|
||||
) -> str | None:
|
||||
@@ -229,6 +265,9 @@ async def call_provider_with_retry(
|
||||
retries = config_manager.get_llm_retries()
|
||||
backoff = config_manager.get_llm_backoff()
|
||||
|
||||
# 检查流式调用配置
|
||||
enable_streaming_llm_call = config_manager.get_enable_streaming_llm_call()
|
||||
|
||||
last_exc = None
|
||||
for attempt in range(1, retries + 1):
|
||||
try:
|
||||
@@ -285,10 +324,16 @@ async def call_provider_with_retry(
|
||||
if extra_generate_kwargs:
|
||||
llm_kwargs.update(extra_generate_kwargs)
|
||||
|
||||
if enable_streaming_llm_call:
|
||||
logger.info("[LLM 调用] 使用流式 Provider 调用")
|
||||
|
||||
async def _invoke_llm(pid: str):
|
||||
if enable_streaming_llm_call:
|
||||
return await _call_provider_stream(context, pid, llm_kwargs)
|
||||
return await context.llm_generate(**llm_kwargs)
|
||||
|
||||
try:
|
||||
llm_resp = await context.llm_generate(
|
||||
**llm_kwargs,
|
||||
)
|
||||
llm_resp = await _invoke_llm(provider_id)
|
||||
except Exception as e:
|
||||
if (
|
||||
response_format is not None
|
||||
@@ -299,9 +344,7 @@ async def call_provider_with_retry(
|
||||
"已自动降级为无 schema 约束重试本次请求。"
|
||||
)
|
||||
llm_kwargs.pop("response_format", None)
|
||||
llm_resp = await context.llm_generate(
|
||||
**llm_kwargs,
|
||||
)
|
||||
llm_resp = await _invoke_llm(provider_id)
|
||||
else:
|
||||
raise
|
||||
|
||||
|
||||
@@ -194,6 +194,10 @@ class ConfigManager:
|
||||
"""获取LLM请求重试退避基值(秒),实际退避会乘以尝试次数"""
|
||||
return self._get_group("llm").get("llm_backoff", 2)
|
||||
|
||||
def get_enable_streaming_llm_call(self) -> bool:
|
||||
"""获取是否启用流式 LLM 调用"""
|
||||
return self._get_group("llm").get("enable_streaming_llm_call", False)
|
||||
|
||||
def get_debug_mode(self) -> bool:
|
||||
"""获取是否启用调试模式"""
|
||||
return self._get_group("basic").get("debug_mode", False)
|
||||
|
||||
Reference in New Issue
Block a user