feat: add inline message editing and regeneration functionality for webui (#7673)

* feat: add inline message editing and regeneration functionality for webui

- Implemented inline editing for user messages in the chat component.
- Added a regenerate menu for retrying messages with different models.
- Enhanced message handling to include llm_checkpoint_id for better tracking.
- Updated localization files to include new actions for retrying and model selection.
- Introduced tests for checkpoint message handling and chat route functionality.

* feat: thread mode in webui

* feat: enhance message editing functionality to allow only the latest user message to be edited

* feat: add error handling and user feedback for thread creation in chat component

* feat: add thread count display and localization support in chat component

* feat: add RefsSidebar component and integrate reference management in chat UI

* feat: improve message editing validation and cleanup for bot messages

* feat: enhance checkpoint message handling with binding and dumping functionality
# Conflicts:
#	astrbot/core/agent/message.py
#	astrbot/core/agent/runners/tool_loop_agent_runner.py
#	astrbot/core/astr_main_agent.py
#	astrbot/core/db/sqlite.py
#	astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py
#	astrbot/core/platform/sources/webchat/webchat_adapter.py
#	astrbot/dashboard/routes/chat.py
#	astrbot/dashboard/routes/live_chat.py
#	dashboard/src/components/chat/Chat.vue
#	dashboard/src/components/chat/ChatInput.vue
#	dashboard/src/components/chat/ChatMessageList.vue
#	dashboard/src/components/chat/ProviderModelMenu.vue
#	dashboard/src/components/shared/StyledMenu.vue
#	dashboard/src/composables/useMessages.ts
#	dashboard/src/i18n/locales/ru-RU/features/chat.json
This commit is contained in:
LIghtJUNction
2026-04-29 03:10:03 +08:00
parent d57582e599
commit 716e651d80
15 changed files with 4005 additions and 4152 deletions
+30 -34
View File
@@ -1,8 +1,7 @@
from __future__ import annotations
# Inspired by MoonshotAI/kosong, credits to MoonshotAI/kosong authors for the original implementation.
# License: Apache License 2.0
from typing import Any, ClassVar, Literal, TypeGuard
from typing import Any, ClassVar, Literal, cast
from pydantic import (
BaseModel,
@@ -15,14 +14,10 @@ from pydantic import (
from pydantic_core import core_schema
def _is_str_keyed_dict(value: object) -> TypeGuard[dict[str, object]]:
return isinstance(value, dict) and all(isinstance(key, str) for key in value)
class ContentPart(BaseModel):
"""A part of the content in a message."""
__content_part_registry: ClassVar[dict[str, type[ContentPart]]] = {}
__content_part_registry: ClassVar[dict[str, type["ContentPart"]]] = {}
type: Literal["text", "think", "image_url", "audio_url"]
@@ -39,25 +34,23 @@ class ContentPart(BaseModel):
@classmethod
def __get_pydantic_core_schema__(
cls,
source_type: object,
handler: GetCoreSchemaHandler,
cls, source_type: Any, handler: GetCoreSchemaHandler
) -> core_schema.CoreSchema:
# If we're dealing with the base ContentPart class, use custom validation
if cls.__name__ == "ContentPart":
def validate_content_part(value: object) -> ContentPart:
def validate_content_part(value: Any) -> Any:
# if it's already an instance of a ContentPart subclass, return it
if isinstance(value, cls):
if hasattr(value, "__class__") and issubclass(value.__class__, cls):
return value
# if it's a dict with a type field, dispatch to the appropriate subclass
if _is_str_keyed_dict(value):
type_value = value.get("type")
if isinstance(type_value, str):
target_class = cls.__content_part_registry.get(type_value)
if target_class is not None:
return target_class.model_validate(value)
if isinstance(value, dict) and "type" in value:
type_value: Any | None = cast(dict[str, Any], value).get("type")
if not isinstance(type_value, str):
raise ValueError(f"Cannot validate {value} as ContentPart")
target_class = cls.__content_part_registry[type_value]
return target_class.model_validate(value)
raise ValueError(f"Cannot validate {value} as ContentPart")
@@ -68,25 +61,27 @@ class ContentPart(BaseModel):
class TextPart(ContentPart):
""">>> TextPart(text="Hello, world!").model_dump()
"""
>>> TextPart(text="Hello, world!").model_dump()
{'type': 'text', 'text': 'Hello, world!'}
"""
type: Literal["text"] = "text"
type: str = "text"
text: str
class ThinkPart(ContentPart):
""">>> ThinkPart(think="I think I need to think about this.").model_dump()
"""
>>> ThinkPart(think="I think I need to think about this.").model_dump()
{'type': 'think', 'think': 'I think I need to think about this.', 'encrypted': None}
"""
type: Literal["think"] = "think"
type: str = "think"
think: str
encrypted: str | None = None
"""Encrypted thinking content, or signature."""
def merge_in_place(self, other: object) -> bool:
def merge_in_place(self, other: Any) -> bool:
if not isinstance(other, ThinkPart):
return False
if self.encrypted:
@@ -98,7 +93,8 @@ class ThinkPart(ContentPart):
class ImageURLPart(ContentPart):
""">>> ImageURLPart(image_url="http://example.com/image.jpg").model_dump()
"""
>>> ImageURLPart(image_url="http://example.com/image.jpg").model_dump()
{'type': 'image_url', 'image_url': 'http://example.com/image.jpg'}
"""
@@ -108,12 +104,13 @@ class ImageURLPart(ContentPart):
id: str | None = None
"""The ID of the image, to allow LLMs to distinguish different images."""
type: Literal["image_url"] = "image_url"
type: str = "image_url"
image_url: ImageURL
class AudioURLPart(ContentPart):
""">>> AudioURLPart(audio_url=AudioURLPart.AudioURL(url="https://example.com/audio.mp3")).model_dump()
"""
>>> AudioURLPart(audio_url=AudioURLPart.AudioURL(url="https://example.com/audio.mp3")).model_dump()
{'type': 'audio_url', 'audio_url': {'url': 'https://example.com/audio.mp3', 'id': None}}
"""
@@ -123,12 +120,13 @@ class AudioURLPart(ContentPart):
id: str | None = None
"""The ID of the audio, to allow LLMs to distinguish different audios."""
type: Literal["audio_url"] = "audio_url"
type: str = "audio_url"
audio_url: AudioURL
class ToolCall(BaseModel):
"""A tool call requested by the assistant.
"""
A tool call requested by the assistant.
>>> ToolCall(
... id="123",
@@ -150,7 +148,7 @@ class ToolCall(BaseModel):
"""The ID of the tool call."""
function: FunctionBody
"""The function body of the tool call."""
extra_content: dict[str, object] | None = None
extra_content: dict[str, Any] | None = None
"""Extra metadata for the tool call."""
@model_serializer(mode="wrap")
@@ -217,7 +215,7 @@ class Message(BaseModel):
# other all cases: content is required
if self.content is None:
raise ValueError(
"content is required unless role='assistant' and tool_calls is not None",
"content is required unless role='assistant' and tool_calls is not None"
)
return self
@@ -334,8 +332,6 @@ def dump_messages_with_checkpoints(messages: list[Message]) -> list[dict]:
dumped.append(message.model_dump())
if message._checkpoint_after is not None:
dumped.append(
CheckpointMessageSegment(
content=message._checkpoint_after,
).model_dump(),
CheckpointMessageSegment(content=message._checkpoint_after).model_dump()
)
return dumped
@@ -1,6 +1,5 @@
import asyncio
import copy
import os
import sys
import time
import traceback
@@ -9,6 +8,7 @@ import uuid
from collections.abc import AsyncIterator
from contextlib import suppress
from dataclasses import dataclass, field, replace
from pathlib import Path
from mcp.types import (
BlobResourceContents,
@@ -26,24 +26,8 @@ from tenacity import (
)
from astrbot import logger
from astrbot.core.agent.context.config import ContextConfig
from astrbot.core.agent.context.manager import ContextManager
from astrbot.core.agent.context.token_counter import EstimateTokenCounter
from astrbot.core.agent.hooks import BaseAgentRunHooks
from astrbot.core.agent.message import (
AssistantMessageSegment,
ImageURLPart,
Message,
TextPart,
ThinkPart,
ToolCallMessageSegment,
bind_checkpoint_messages,
)
from astrbot.core.agent.response import AgentResponse, AgentResponseData, AgentStats
from astrbot.core.agent.run_context import ContextWrapper, TContext
from astrbot.core.agent.runners.base import AgentState, BaseAgentRunner
from astrbot.core.agent.message import ImageURLPart, TextPart, ThinkPart
from astrbot.core.agent.tool import FunctionTool, ToolSet
from astrbot.core.agent.tool_executor import BaseFunctionToolExecutor
from astrbot.core.agent.tool_image_cache import tool_image_cache
from astrbot.core.exceptions import EmptyModelOutputError
from astrbot.core.message.components import Json
@@ -64,6 +48,22 @@ from astrbot.core.provider.modalities import (
)
from astrbot.core.provider.provider import Provider
from ..context.compressor import ContextCompressor
from ..context.config import ContextConfig
from ..context.manager import ContextManager
from ..context.token_counter import EstimateTokenCounter, TokenCounter
from ..hooks import BaseAgentRunHooks
from ..message import (
AssistantMessageSegment,
Message,
ToolCallMessageSegment,
bind_checkpoint_messages,
)
from ..response import AgentResponseData, AgentStats
from ..run_context import ContextWrapper, TContext
from ..tool_executor import BaseFunctionToolExecutor
from .base import AgentResponse, AgentState, BaseAgentRunner
if sys.version_info >= (3, 12):
from typing import override
else:
@@ -83,8 +83,7 @@ class _HandleFunctionToolsResult:
@classmethod
def from_tool_call_result_blocks(
cls,
blocks: list[ToolCallMessageSegment],
cls, blocks: list[ToolCallMessageSegment]
) -> "_HandleFunctionToolsResult":
return cls(kind="tool_call_result_blocks", tool_call_result_blocks=blocks)
@@ -184,12 +183,12 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
self.stats.end_time = time.time()
parts = []
if llm_resp.reasoning_content is not None or llm_resp.reasoning_signature:
if llm_resp.reasoning_content or llm_resp.reasoning_signature:
parts.append(
ThinkPart(
think=llm_resp.reasoning_content or "",
think=llm_resp.reasoning_content,
encrypted=llm_resp.reasoning_signature,
),
)
)
if llm_resp.completion_text:
parts.append(TextPart(text=llm_resp.completion_text))
@@ -222,21 +221,14 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
# truncate by turns compressor
truncate_turns: int = 1,
# customize
custom_token_counter: T.Any = None,
custom_compressor: T.Any = None,
custom_token_counter: TokenCounter | None = None,
custom_compressor: ContextCompressor | None = None,
tool_schema_mode: str | None = "full",
fallback_providers: list[Provider] | None = None,
provider_config: dict | None = None,
tool_result_overflow_dir: str | None = None,
read_tool: FunctionTool | None = None,
**kwargs: T.Any,
) -> None:
raw_tool_result_overflow_dir = kwargs.get("tool_result_overflow_dir")
tool_result_overflow_dir = (
raw_tool_result_overflow_dir
if isinstance(raw_tool_result_overflow_dir, str)
else None
)
raw_read_tool = kwargs.get("read_tool")
read_tool = raw_read_tool if isinstance(raw_read_tool, FunctionTool) else None
self.req = request
self.streaming = streaming
self.enforce_max_turns = enforce_max_turns
@@ -380,7 +372,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
if self.tool_result_overflow_dir is None:
raise ValueError("tool_result_overflow_dir is not configured")
overflow_dir = self.tool_result_overflow_dir
overflow_dir = Path(self.tool_result_overflow_dir).resolve(strict=False)
safe_tool_call_id = (
"".join(
ch if ch.isalnum() or ch in {"-", "_", "."} else "_"
@@ -389,14 +381,12 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
or "tool_call"
)
file_name = f"{safe_tool_call_id}_{uuid.uuid4().hex[:8]}.txt"
overflow_path = overflow_dir / file_name
def _run() -> str:
normalized_overflow_dir = os.path.abspath(overflow_dir)
overflow_path = os.path.join(normalized_overflow_dir, file_name)
os.makedirs(normalized_overflow_dir, exist_ok=True)
with open(overflow_path, "w", encoding="utf-8") as file:
file.write(content)
return overflow_path
overflow_dir.mkdir(parents=True, exist_ok=True)
overflow_path.write_text(content, encoding="utf-8")
return str(overflow_path)
return await asyncio.to_thread(_run)
@@ -410,7 +400,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
return content
estimated_tokens = self._tool_result_token_counter.count_tokens(
[Message(role="tool", content=content, tool_call_id=tool_call_id)],
[Message(role="tool", content=content, tool_call_id=tool_call_id)]
)
if estimated_tokens <= self.TOOL_RESULT_MAX_ESTIMATED_TOKENS:
return content
@@ -455,7 +445,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
preview = content
while preview:
estimated_tokens = self._tool_result_token_counter.count_tokens(
[Message(role="tool", content=preview, tool_call_id=tool_call_id)],
[Message(role="tool", content=preview, tool_call_id=tool_call_id)]
)
if estimated_tokens <= self.TOOL_RESULT_PREVIEW_MAX_ESTIMATED_TOKENS:
return preview
@@ -466,34 +456,25 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
return preview
async def _iter_llm_responses(
self,
*,
include_model: bool = True,
self, *, include_model: bool = True
) -> T.AsyncGenerator[LLMResponse, None]:
"""Yields chunks *and* a final LLMResponse."""
contexts = self._sanitize_contexts_for_provider(self.run_context.messages)
func_tool = self._func_tool_for_provider()
model = self.req.model if include_model else None
payload = {
"contexts": self._sanitize_contexts_for_provider(self.run_context.messages),
"func_tool": self._func_tool_for_provider(),
"session_id": self.req.session_id,
"extra_user_content_parts": self.req.extra_user_content_parts, # list[ContentPart]
"abort_signal": self._abort_signal,
}
if include_model:
# For primary provider we keep explicit model selection if provided.
payload["model"] = self.req.model
if self.streaming:
stream = self.provider.text_chat_stream(
contexts=contexts,
func_tool=func_tool,
session_id=self.req.session_id,
extra_user_content_parts=self.req.extra_user_content_parts,
model=model,
abort_signal=self._abort_signal,
)
async for resp in stream:
stream = self.provider.text_chat_stream(**payload)
async for resp in stream: # type: ignore
yield resp
else:
yield await self.provider.text_chat(
contexts=contexts,
func_tool=func_tool,
session_id=self.req.session_id,
extra_user_content_parts=self.req.extra_user_content_parts,
model=model,
abort_signal=self._abort_signal,
)
yield await self.provider.text_chat(**payload)
async def _iter_llm_responses_with_fallback(
self,
@@ -531,7 +512,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
with attempt:
try:
async for resp in self._iter_llm_responses(
include_model=idx == 0,
include_model=idx == 0
):
if resp.is_chunk:
has_stream_output = True
@@ -727,25 +708,16 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
token_usage = self.req.conversation.token_usage if self.req.conversation else 0
self._simple_print_message_role("[BefCompact]")
self.run_context.messages = await self.context_manager.process(
self.run_context.messages,
trusted_token_usage=token_usage,
self.run_context.messages, trusted_token_usage=token_usage
)
self._simple_print_message_role("[AftCompact]")
async for llm_response in self._iter_llm_responses_with_fallback():
if llm_response.is_chunk:
# update ttft
if self.stats.time_to_first_token == 0:
self.stats.time_to_first_token = time.time() - self.stats.start_time
if llm_response.reasoning_content:
yield AgentResponse(
type="streaming_delta",
data=AgentResponseData(
chain=MessageChain(type="reasoning").message(
llm_response.reasoning_content,
),
),
)
if llm_response.result_chain:
yield AgentResponse(
type="streaming_delta",
@@ -758,6 +730,15 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
chain=MessageChain().message(llm_response.completion_text),
),
)
elif llm_response.reasoning_content:
yield AgentResponse(
type="streaming_delta",
data=AgentResponseData(
chain=MessageChain(type="reasoning").message(
llm_response.reasoning_content,
),
),
)
if self._is_stop_requested():
llm_resp_result = LLMResponse(
role="assistant",
@@ -811,15 +792,6 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
await self._complete_with_assistant_response(llm_resp)
# 返回 LLM 结果
if llm_resp.reasoning_content:
yield AgentResponse(
type="llm_result",
data=AgentResponseData(
chain=MessageChain(type="reasoning").message(
llm_resp.reasoning_content,
),
),
)
if llm_resp.result_chain:
yield AgentResponse(
type="llm_result",
@@ -839,17 +811,8 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
llm_resp, _ = await self._resolve_tool_exec(llm_resp)
if not llm_resp.tools_call_name:
logger.warning(
"skills_like tool re-query returned no tool calls; fallback to assistant response.",
"skills_like tool re-query returned no tool calls; fallback to assistant response."
)
if llm_resp.reasoning_content:
yield AgentResponse(
type="llm_result",
data=AgentResponseData(
chain=MessageChain(type="reasoning").message(
llm_resp.reasoning_content,
),
),
)
if llm_resp.result_chain:
yield AgentResponse(
type="llm_result",
@@ -862,7 +825,6 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
chain=MessageChain().message(llm_resp.completion_text),
),
)
await self._complete_with_assistant_response(llm_resp)
return
@@ -896,12 +858,12 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
# 将结果添加到上下文中
parts = []
if llm_resp.reasoning_content is not None or llm_resp.reasoning_signature:
if llm_resp.reasoning_content or llm_resp.reasoning_signature:
parts.append(
ThinkPart(
think=llm_resp.reasoning_content or "",
think=llm_resp.reasoning_content,
encrypted=llm_resp.reasoning_signature,
),
)
)
if llm_resp.completion_text:
parts.append(TextPart(text=llm_resp.completion_text))
@@ -916,7 +878,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
)
# record the assistant message with tool calls
self.run_context.messages.extend(
tool_calls_result.to_openai_messages_model(),
tool_calls_result.to_openai_messages_model()
)
# If there are cached images and the model supports image input,
@@ -929,37 +891,35 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
image_parts = []
for cached_img in cached_images:
img_data = tool_image_cache.get_image_base64_by_path(
cached_img.file_path,
cached_img.mime_type,
cached_img.file_path, cached_img.mime_type
)
if img_data:
base64_data, mime_type = img_data
image_parts.append(
TextPart(
text=f"[Image from tool '{cached_img.tool_name}', path='{cached_img.file_path}']",
),
text=f"[Image from tool '{cached_img.tool_name}', path='{cached_img.file_path}']"
)
)
image_parts.append(
ImageURLPart(
image_url=ImageURLPart.ImageURL(
url=f"data:{mime_type};base64,{base64_data}",
id=cached_img.file_path,
),
),
)
)
)
if image_parts:
self.run_context.messages.append(
Message(role="user", content=image_parts),
Message(role="user", content=image_parts)
)
logger.debug(
f"Appended {len(cached_images)} cached image(s) to context for LLM review",
f"Appended {len(cached_images)} cached image(s) to context for LLM review"
)
self.req.append_tool_calls_result(tool_calls_result)
async def step_until_done(
self,
max_step: int,
self, max_step: int
) -> T.AsyncGenerator[AgentResponse, None]:
"""Process steps until the agent is done."""
step_count = 0
@@ -971,7 +931,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
# 如果循环结束了但是 agent 还没有完成,说明是达到了 max_step
if not self.done():
logger.warning(
f"Agent reached max steps ({max_step}), forcing a final response.",
f"Agent reached max steps ({max_step}), forcing a final response."
)
# 拔掉所有工具
if self.req:
@@ -981,7 +941,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
Message(
role="user",
content=self.MAX_STEPS_REACHED_PROMPT,
),
)
)
# 再执行最后一步
async for resp in self.step():
@@ -1010,9 +970,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
llm_response.tools_call_name,
llm_response.tools_call_args,
llm_response.tools_call_ids,
strict=False,
):
tool_result_blocks_start = len(tool_call_result_blocks)
tool_call_streak = self._track_tool_call_streak(func_tool_name)
yield _HandleFunctionToolsResult.from_message_chain(
MessageChain(
@@ -1024,10 +982,10 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
"name": func_tool_name,
"args": func_tool_args,
"ts": time.time(),
},
),
}
)
],
),
)
)
try:
if not req.func_tool:
@@ -1126,11 +1084,11 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
result_parts.append(
f"Image returned and cached at path='{cached_img.file_path}'. "
f"Review the image below. Use send_message_to_user to send it to the user if satisfied, "
f"with type='image' and path='{cached_img.file_path}'.",
f"with type='image' and path='{cached_img.file_path}'."
)
# Yield image info for LLM visibility (will be handled in step())
yield _HandleFunctionToolsResult.from_cached_image(
cached_img,
cached_img
)
elif isinstance(content_item, EmbeddedResource):
resource = content_item.resource
@@ -1152,15 +1110,15 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
result_parts.append(
f"Image returned and cached at path='{cached_img.file_path}'. "
f"Review the image below. Use send_message_to_user to send it to the user if satisfied, "
f"with type='image' and path='{cached_img.file_path}'.",
f"with type='image' and path='{cached_img.file_path}'."
)
# Yield image info for LLM visibility
yield _HandleFunctionToolsResult.from_cached_image(
cached_img,
cached_img
)
else:
result_parts.append(
"The tool has returned a data type that is not supported.",
"The tool has returned a data type that is not supported."
)
if result_parts:
inline_result = "\n\n".join(result_parts)
@@ -1172,8 +1130,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
func_tool_id,
inline_result
+ self._build_repeated_tool_call_guidance(
func_tool_name,
tool_call_streak,
func_tool_name, tool_call_streak
),
)
@@ -1182,7 +1139,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
# 这里我们将直接结束 Agent Loop
# 发送消息逻辑在 ToolExecutor 中处理了
logger.warning(
f"{func_tool_name} 没有返回值,或者已将结果直接发送给用户。",
f"{func_tool_name} 没有返回值,或者已将结果直接发送给用户。"
)
self._transition_state(AgentState.DONE)
self.stats.end_time = time.time()
@@ -1190,8 +1147,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
func_tool_id,
"The tool has no return value, or has sent the result directly to the user."
+ self._build_repeated_tool_call_guidance(
func_tool_name,
tool_call_streak,
func_tool_name, tool_call_streak
),
)
else:
@@ -1203,8 +1159,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
func_tool_id,
"*The tool has returned an unsupported type. Please tell the user to check the definition and implementation of this tool.*"
+ self._build_repeated_tool_call_guidance(
func_tool_name,
tool_call_streak,
func_tool_name, tool_call_streak
),
)
@@ -1225,33 +1180,33 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
func_tool_id,
f"error: {e!s}"
+ self._build_repeated_tool_call_guidance(
func_tool_name,
tool_call_streak,
func_tool_name, tool_call_streak
),
)
if len(tool_call_result_blocks) > tool_result_blocks_start:
tool_result_content = str(tool_call_result_blocks[-1].content)
yield _HandleFunctionToolsResult.from_message_chain(
MessageChain(
type="tool_call_result",
chain=[
Json(
data={
"id": func_tool_id,
"ts": time.time(),
"result": tool_result_content,
}
)
],
)
# yield the last tool call result
if tool_call_result_blocks:
last_tcr_content = str(tool_call_result_blocks[-1].content)
yield _HandleFunctionToolsResult.from_message_chain(
MessageChain(
type="tool_call_result",
chain=[
Json(
data={
"id": func_tool_id,
"ts": time.time(),
"result": last_tcr_content,
}
)
],
)
logger.info(f"Tool `{func_tool_name}` Result: {tool_result_content}")
)
logger.info(f"Tool `{func_tool_name}` Result: {last_tcr_content}")
# 处理函数调用响应
if tool_call_result_blocks:
yield _HandleFunctionToolsResult.from_tool_call_result_blocks(
tool_call_result_blocks,
tool_call_result_blocks
)
def _build_tool_requery_context(
@@ -1267,7 +1222,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
elif isinstance(msg, dict):
contexts.append(copy.deepcopy(msg))
instruction = self.SKILLS_LIKE_REQUERY_INSTRUCTION_TEMPLATE.format(
tool_names=", ".join(tool_names),
tool_names=", ".join(tool_names)
)
if extra_instruction:
instruction = f"{instruction}\n{extra_instruction}"
@@ -1310,8 +1265,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
if isinstance(self._tool_schema_param_set, ToolSet):
param_subset = self._build_tool_subset(
self._tool_schema_param_set,
tool_names,
self._tool_schema_param_set, tool_names
)
if param_subset.tools and tool_names:
contexts = self._build_tool_requery_context(tool_names)
@@ -1321,7 +1275,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
model=self.req.model,
session_id=self.req.session_id,
extra_user_content_parts=self.req.extra_user_content_parts,
# tool_choice="required",
tool_choice="required",
abort_signal=self._abort_signal,
)
if requery_resp:
@@ -1335,7 +1289,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
and not self._has_meaningful_assistant_reply(llm_resp)
):
logger.warning(
"skills_like tool re-query returned no tool calls and no explanation; retrying with stronger instruction.",
"skills_like tool re-query returned no tool calls and no explanation; retrying with stronger instruction."
)
repair_contexts = self._build_tool_requery_context(
tool_names,
@@ -1347,7 +1301,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
model=self.req.model,
session_id=self.req.session_id,
extra_user_content_parts=self.req.extra_user_content_parts,
# tool_choice="required",
tool_choice="required",
abort_signal=self._abort_signal,
)
if repair_resp:
@@ -1389,12 +1343,12 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
self.stats.end_time = time.time()
parts = []
if llm_resp.reasoning_content is not None or llm_resp.reasoning_signature:
if llm_resp.reasoning_content or llm_resp.reasoning_signature:
parts.append(
ThinkPart(
think=llm_resp.reasoning_content or "",
think=llm_resp.reasoning_content,
encrypted=llm_resp.reasoning_signature,
),
)
)
if llm_resp.completion_text:
parts.append(TextPart(text=llm_resp.completion_text))
@@ -1423,17 +1377,14 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
self,
executor: AsyncIterator[ToolExecutorResultT],
) -> T.AsyncGenerator[ToolExecutorResultT, None]:
async def _get_next_result() -> ToolExecutorResultT:
return await anext(executor)
while True:
if self._is_stop_requested():
await self._close_executor(executor)
raise _ToolExecutionInterrupted(
"Tool execution interrupted before reading the next tool result.",
"Tool execution interrupted before reading the next tool result."
)
next_result_task = asyncio.create_task(_get_next_result())
next_result_task = asyncio.create_task(anext(executor))
abort_task = asyncio.create_task(self._abort_signal.wait())
try:
done, _ = await asyncio.wait(
@@ -1450,7 +1401,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
await self._close_executor(executor)
raise _ToolExecutionInterrupted(
"Tool execution interrupted by a stop request.",
"Tool execution interrupted by a stop request."
)
try:
+48 -91
View File
@@ -10,7 +10,6 @@ import zoneinfo
from collections.abc import Coroutine
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from astrbot.core import logger
from astrbot.core.agent.handoff import HandoffTool
@@ -48,9 +47,6 @@ from astrbot.core.tools.computer_tools import (
BrowserExecTool,
CreateSkillCandidateTool,
CreateSkillPayloadTool,
CuaKeyboardTypeTool,
CuaMouseClickTool,
CuaScreenshotTool,
EvaluateSkillCandidateTool,
ExecuteShellTool,
FileDownloadTool,
@@ -76,12 +72,11 @@ from astrbot.core.tools.knowledge_base_tools import (
KnowledgeBaseQueryTool,
retrieve_knowledge_base,
)
from astrbot.core.tools.message_tools import SendMessageToUserTool
from astrbot.core.tools.web_search_tools import (
BaiduWebSearchTool,
BochaWebSearchTool,
BraveWebSearchTool,
FirecrawlExtractWebPageTool,
FirecrawlWebSearchTool,
TavilyExtractWebPageTool,
TavilyWebSearchTool,
normalize_legacy_web_search_config,
@@ -113,8 +108,7 @@ from astrbot.core.utils.string_utils import normalize_and_dedupe_strings
@dataclass(slots=True)
class MainAgentBuildConfig:
"""The main agent build configuration.
Most of the configs can be found in the cmd_config.json
"""
Most of the configs can be found in the cmd_config.json"""
tool_call_timeout: int
"""The timeout (in seconds) for a tool call.
@@ -178,8 +172,7 @@ class MainAgentBuildResult:
def _select_provider(
event: AstrMessageEvent,
plugin_context: Context,
event: AstrMessageEvent, plugin_context: Context
) -> Provider | None:
"""Select chat provider for the event."""
sel_provider = event.get_extra("selected_provider")
@@ -189,8 +182,7 @@ def _select_provider(
logger.error("未找到指定的提供商: %s。", sel_provider)
if not isinstance(provider, Provider):
logger.error(
"选择的提供商类型无效(%s),跳过 LLM 请求处理。",
type(provider),
"选择的提供商类型无效(%s),跳过 LLM 请求处理。", type(provider)
)
return None
return provider
@@ -202,8 +194,7 @@ def _select_provider(
async def _get_session_conv(
event: AstrMessageEvent,
plugin_context: Context,
event: AstrMessageEvent, plugin_context: Context
) -> Conversation:
conv_mgr = plugin_context.conversation_manager
umo = event.unified_msg_origin
@@ -247,8 +238,8 @@ async def _apply_kb(
req.func_tool = ToolSet()
req.func_tool.add_tool(
plugin_context.get_llm_tool_manager().get_builtin_tool(
KnowledgeBaseQueryTool,
),
KnowledgeBaseQueryTool
)
)
@@ -283,13 +274,13 @@ async def _apply_file_extract(
config.file_extract_msh_api_key,
)
for file_path in file_paths
],
]
)
else:
logger.error("Unsupported file extract provider: %s", config.file_extract_prov)
return
for file_content, file_name in zip(file_contents, file_names, strict=False):
for file_content, file_name in zip(file_contents, file_names):
req.contexts.append(
{
"role": "system",
@@ -400,8 +391,7 @@ async def _ensure_persona_and_skills(
)
set_persona_custom_error_message_on_event(
event,
extract_persona_custom_error_message_from_persona(persona),
event, extract_persona_custom_error_message_from_persona(persona)
)
if persona:
@@ -482,7 +472,7 @@ async def _ensure_persona_and_skills(
tool.name
for tool in tmgr.func_list
if not isinstance(tool, HandoffTool)
],
]
)
continue
if not isinstance(tools, list):
@@ -574,7 +564,7 @@ async def _ensure_img_caption(
)
if caption:
req.extra_user_content_parts.append(
TextPart(text=f"<image_caption>{caption}</image_caption>"),
TextPart(text=f"<image_caption>{caption}</image_caption>")
)
req.image_urls = []
except Exception as exc: # noqa: BLE001
@@ -586,19 +576,19 @@ async def _ensure_img_caption(
def _append_quoted_image_attachment(req: ProviderRequest, image_path: str) -> None:
req.extra_user_content_parts.append(
TextPart(text=f"[Image Attachment in quoted message: path {image_path}]"),
TextPart(text=f"[Image Attachment in quoted message: path {image_path}]")
)
def _append_audio_attachment(req: ProviderRequest, audio_path: str) -> None:
req.extra_user_content_parts.append(
TextPart(text=f"[Audio Attachment: path {audio_path}]"),
TextPart(text=f"[Audio Attachment: path {audio_path}]")
)
def _append_quoted_audio_attachment(req: ProviderRequest, audio_path: str) -> None:
req.extra_user_content_parts.append(
TextPart(text=f"[Audio Attachment in quoted message: path {audio_path}]"),
TextPart(text=f"[Audio Attachment in quoted message: path {audio_path}]")
)
@@ -637,7 +627,7 @@ def _get_quoted_message_parser_settings(
overrides = provider_settings.get("quoted_message_parser")
if not isinstance(overrides, dict):
return DEFAULT_QUOTED_MESSAGE_SETTINGS
return DEFAULT_QUOTED_MESSAGE_SETTINGS.with_overrides(overrides) # type: ignore
return DEFAULT_QUOTED_MESSAGE_SETTINGS.with_overrides(overrides)
def _get_image_compress_args(
@@ -650,10 +640,8 @@ def _get_image_compress_args(
if not isinstance(enabled, bool):
enabled = True
raw_options = provider_settings.get("image_compress_options")
options: dict[str, Any] = {}
if isinstance(raw_options, dict):
options = {str(k): v for k, v in raw_options.items()}
raw_options = provider_settings.get("image_compress_options", {})
options = raw_options if isinstance(raw_options, dict) else {}
max_size = options.get("max_size", IMAGE_COMPRESS_DEFAULT_MAX_SIZE)
if not isinstance(max_size, int):
@@ -753,25 +741,22 @@ async def _process_quote_message(
)
if llm_resp.completion_text:
content_parts.append(
f"[Image Caption in quoted message]: {llm_resp.completion_text}",
f"[Image Caption in quoted message]: {llm_resp.completion_text}"
)
else:
logger.warning("No provider found for image captioning in quote.")
except BaseException as exc:
logger.error("处理引用图片失败: %s", exc)
finally:
if compress_path and compress_path != path:
from anyio import Path as AnyioPath
compress_file = AnyioPath(compress_path)
if await compress_file.exists():
try:
await compress_file.unlink()
except Exception as exc: # noqa: BLE001
logger.warning(
"Fail to remove temporary compressed image: %s",
exc,
)
if (
compress_path
and compress_path != path
and os.path.exists(compress_path)
):
try:
os.remove(compress_path)
except Exception as exc: # noqa: BLE001
logger.warning("Fail to remove temporary compressed image: %s", exc)
quoted_content = "\n".join(content_parts)
quoted_text = f"<Quoted Message>\n{quoted_content}\n</Quoted Message>"
@@ -829,7 +814,7 @@ async def _decorate_llm_request(
config: MainAgentBuildConfig,
) -> None:
cfg = config.provider_settings or plugin_context.get_config(
umo=event.unified_msg_origin,
umo=event.unified_msg_origin
).get("provider_settings", {})
_apply_prompt_prefix(req, cfg)
@@ -878,7 +863,7 @@ def _plugin_tool_fix(event: AstrMessageEvent, req: ProviderRequest) -> None:
# 保留 MCP 工具
new_tool_set.add_tool(tool)
continue
mp = getattr(tool, "handler_module_path", None)
mp = tool.handler_module_path
if not mp:
# 没有 plugin 归属信息的工具(如 subagent transfer_to_*)
# 不应受到会话插件过滤影响。
@@ -895,9 +880,7 @@ def _plugin_tool_fix(event: AstrMessageEvent, req: ProviderRequest) -> None:
async def _handle_webchat(
event: AstrMessageEvent,
req: ProviderRequest,
prov: Provider,
event: AstrMessageEvent, req: ProviderRequest, prov: Provider
) -> None:
from astrbot.core import db_helper
@@ -932,9 +915,7 @@ async def _handle_webchat(
if not title or "<None>" in title:
return
logger.info(
"Generated chatui title for session %s: %s",
chatui_session_id,
title,
"Generated chatui title for session %s: %s", chatui_session_id, title
)
await db_helper.update_platform_session(
session_id=chatui_session_id,
@@ -1032,22 +1013,6 @@ def _apply_sandbox_tools(
req.func_tool.add_tool(tool_mgr.get_builtin_tool(RollbackSkillReleaseTool))
req.func_tool.add_tool(tool_mgr.get_builtin_tool(SyncSkillReleaseTool))
if booter == "cua":
req.system_prompt += (
"\n[CUA Desktop Control]\n"
"Use `astrbot_execute_shell` with `background=true` to launch GUI apps. "
'Use Firefox for browser tasks, for example `firefox "https://example.com"`. '
"After each visible step, call `astrbot_cua_screenshot` with "
"`send_to_user=true` and `return_image_to_llm=true` so the user can "
"monitor progress. When typing, inspect the screenshot first and confirm "
"the target field is focused and empty or safe to append to. Use "
"`astrbot_cua_mouse_click` for coordinates and `astrbot_cua_keyboard_type` "
"for text input; use text=`\\n` for Enter.\n"
)
req.func_tool.add_tool(tool_mgr.get_builtin_tool(CuaScreenshotTool))
req.func_tool.add_tool(tool_mgr.get_builtin_tool(CuaMouseClickTool))
req.func_tool.add_tool(tool_mgr.get_builtin_tool(CuaKeyboardTypeTool))
req.system_prompt = f"{req.system_prompt or ''}\n{SANDBOX_MODE_PROMPT}\n"
@@ -1082,16 +1047,12 @@ async def _apply_web_search_tools(
req.func_tool.add_tool(tool_mgr.get_builtin_tool(BochaWebSearchTool))
elif provider == "brave":
req.func_tool.add_tool(tool_mgr.get_builtin_tool(BraveWebSearchTool))
elif provider == "firecrawl":
req.func_tool.add_tool(tool_mgr.get_builtin_tool(FirecrawlWebSearchTool))
req.func_tool.add_tool(tool_mgr.get_builtin_tool(FirecrawlExtractWebPageTool))
elif provider == "baidu_ai_search":
req.func_tool.add_tool(tool_mgr.get_builtin_tool(BaiduWebSearchTool))
def _get_compress_provider(
config: MainAgentBuildConfig,
plugin_context: Context,
config: MainAgentBuildConfig, plugin_context: Context
) -> Provider | None:
if not config.llm_compress_provider_id:
return None
@@ -1114,14 +1075,12 @@ def _get_compress_provider(
def _get_fallback_chat_providers(
provider: Provider,
plugin_context: Context,
provider_settings: dict,
provider: Provider, plugin_context: Context, provider_settings: dict
) -> list[Provider]:
fallback_ids = provider_settings.get("fallback_chat_models", [])
if not isinstance(fallback_ids, list):
logger.warning(
"fallback_chat_models setting is not a list, skip fallback providers.",
"fallback_chat_models setting is not a list, skip fallback providers."
)
return []
@@ -1184,7 +1143,7 @@ async def build_main_agent(
if sel_model := event.get_extra("selected_model"):
req.model = sel_model
if config.provider_wake_prefix and not event.message_str.startswith(
config.provider_wake_prefix,
config.provider_wake_prefix
):
return None
@@ -1202,7 +1161,7 @@ async def build_main_agent(
event.track_temporary_local_file(image_path)
req.image_urls.append(image_path)
req.extra_user_content_parts.append(
TextPart(text=f"[Image Attachment: path {image_path}]"),
TextPart(text=f"[Image Attachment: path {image_path}]")
)
elif isinstance(comp, Record):
audio_path = await comp.convert_to_file_path()
@@ -1213,8 +1172,8 @@ async def build_main_agent(
file_name = comp.name or os.path.basename(file_path)
req.extra_user_content_parts.append(
TextPart(
text=f"[File Attachment: name {file_name}, path {file_path}]",
),
text=f"[File Attachment: name {file_name}, path {file_path}]"
)
)
elif isinstance(comp, Video):
await _append_video_attachment(req, comp)
@@ -1223,7 +1182,7 @@ async def build_main_agent(
comp for comp in event.message_obj.message if isinstance(comp, Reply)
]
quoted_message_settings = _get_quoted_message_parser_settings(
config.provider_settings,
config.provider_settings
)
fallback_quoted_image_count = 0
for comp in reply_comps:
@@ -1253,8 +1212,8 @@ async def build_main_agent(
text=(
f"[File Attachment in quoted message: "
f"name {file_name}, path {file_path}]"
),
),
)
)
)
elif isinstance(reply_comp, Video):
await _append_video_attachment(req, reply_comp, quoted=True)
@@ -1268,7 +1227,7 @@ async def build_main_agent(
event,
comp,
settings=quoted_message_settings,
),
)
)
remaining_limit = max(
config.max_quoted_fallback_images
@@ -1321,8 +1280,8 @@ async def build_main_agent(
"The user is asking in a side thread about this selected "
"excerpt from the previous assistant answer:\n"
f"<selected_excerpt>{thread_selected_text.strip()}</selected_excerpt>"
),
),
)
)
)
req.image_urls = normalize_and_dedupe_strings(req.image_urls)
req.audio_urls = normalize_and_dedupe_strings(req.audio_urls)
@@ -1371,8 +1330,8 @@ async def build_main_agent(
req.func_tool = ToolSet()
req.func_tool.add_tool(
plugin_context.get_llm_tool_manager().get_builtin_tool(
"send_message_to_user",
),
SendMessageToUserTool
)
)
if provider.provider_config.get("max_context_tokens", 0) <= 0:
@@ -1423,9 +1382,7 @@ async def build_main_agent(
enforce_max_turns=config.max_context_length,
tool_schema_mode=config.tool_schema_mode,
fallback_providers=_get_fallback_chat_providers(
provider,
plugin_context,
config.provider_settings,
provider, plugin_context, config.provider_settings
),
tool_result_overflow_dir=(
get_astrbot_system_tmp_path()
+653 -550
View File
File diff suppressed because it is too large Load Diff
@@ -86,23 +86,19 @@ class InternalAgentSubStage(Stage):
self.file_extract_enabled: bool = file_extract_conf.get("enable", False)
self.file_extract_prov: str = file_extract_conf.get("provider", "moonshotai")
self.file_extract_msh_api_key: str = file_extract_conf.get(
"moonshotai_api_key",
"",
"moonshotai_api_key", ""
)
# 上下文管理相关
self.context_limit_reached_strategy: str = settings.get(
"context_limit_reached_strategy",
"truncate_by_turns",
"context_limit_reached_strategy", "truncate_by_turns"
)
self.llm_compress_instruction: str = settings.get(
"llm_compress_instruction",
"",
"llm_compress_instruction", ""
)
self.llm_compress_keep_recent: int = settings.get("llm_compress_keep_recent", 4)
self.llm_compress_provider_id: str = settings.get(
"llm_compress_provider_id",
"",
"llm_compress_provider_id", ""
)
self.max_context_length = settings["max_context_length"] # int
self.dequeue_context_length: int = min(
@@ -114,8 +110,7 @@ class InternalAgentSubStage(Stage):
self.llm_safety_mode = settings.get("llm_safety_mode", True)
self.safety_mode_strategy = settings.get(
"safety_mode_strategy",
"system_prompt",
"safety_mode_strategy", "system_prompt"
)
self.computer_use_runtime = settings.get("computer_use_runtime")
@@ -153,9 +148,7 @@ class InternalAgentSubStage(Stage):
)
async def process(
self,
event: AstrMessageEvent,
provider_wake_prefix: str,
self, event: AstrMessageEvent, provider_wake_prefix: str
) -> AsyncGenerator[None, None]:
follow_up_capture: FollowUpCapture | None = None
follow_up_consumed_marked = False
@@ -275,13 +268,13 @@ class InternalAgentSubStage(Stage):
# 获取 TTS Provider
tts_provider = (
self.ctx.plugin_manager.context.get_using_tts_provider(
event.unified_msg_origin,
event.unified_msg_origin
)
)
if not tts_provider:
logger.warning(
"[Live Mode] TTS Provider 未配置,将使用普通流式模式",
"[Live Mode] TTS Provider 未配置,将使用普通流式模式"
)
# 使用 run_live_agent,总是使用流式响应
@@ -376,7 +369,7 @@ class InternalAgentSubStage(Stage):
req,
agent_runner,
final_resp,
),
)
)
# 检查事件是否被停止,如果被停止则不保存历史记录
@@ -404,7 +397,7 @@ class InternalAgentSubStage(Stage):
except Exception as e:
logger.error(f"Error occurred while processing agent: {e}")
custom_error_message = extract_persona_custom_error_message_from_event(
event,
event
)
error_text = custom_error_message or (
f"Error occurred while processing agent request: {e}"
@@ -472,7 +465,7 @@ class InternalAgentSubStage(Stage):
message_to_save.append(
CheckpointMessageSegment(
content=CheckpointData(id=checkpoint_id),
).model_dump(),
).model_dump()
)
# if user_aborted:
@@ -17,9 +17,9 @@ from astrbot.core.platform import (
PlatformMetadata,
)
from astrbot.core.platform.astr_message_event import MessageSesion
from astrbot.core.platform.register import register_platform_adapter
from astrbot.core.utils.astrbot_path import get_astrbot_data_path
from ...register import register_platform_adapter
from .message_parts_helper import (
message_chain_to_storage_message_parts,
parse_webchat_message_parts,
@@ -66,11 +66,13 @@ class WebChatAdapter(Platform):
event_queue: asyncio.Queue,
) -> None:
super().__init__(platform_config, event_queue)
self.settings = platform_settings
self.imgs_dir = os.path.join(get_astrbot_data_path(), "webchat", "imgs")
self.attachments_dir = Path(get_astrbot_data_path()) / "attachments"
os.makedirs(self.imgs_dir, exist_ok=True)
self.attachments_dir.mkdir(parents=True, exist_ok=True)
self.metadata = PlatformMetadata(
name="webchat",
description="webchat",
@@ -87,13 +89,16 @@ class WebChatAdapter(Platform):
) -> None:
conversation_id = _extract_conversation_id(session.session_id)
active_request_ids = self._webchat_queue_mgr.list_back_request_ids(
conversation_id,
conversation_id
)
stream_request_ids = [
req_id for req_id in active_request_ids if not req_id.startswith("ws_sub_")
]
target_request_ids = stream_request_ids or active_request_ids
if not target_request_ids:
# No active streams to consume this proactive message.
# Persist directly and return to avoid creating an unused queue.
try:
await self._save_proactive_message(conversation_id, message_chain)
except Exception as e:
@@ -103,6 +108,7 @@ class WebChatAdapter(Platform):
)
await super().send_by_session(session, message_chain)
return
for request_id in target_request_ids:
await WebChatMessageEvent._send(
request_id,
@@ -111,6 +117,10 @@ class WebChatAdapter(Platform):
streaming=True,
emit_complete=True,
)
# If only passive subscription queues exist for this conversation,
# keep a proactive save as a fallback since they are not tied to
# the normal streaming persistence path.
if not stream_request_ids:
try:
await self._save_proactive_message(conversation_id, message_chain)
@@ -119,6 +129,7 @@ class WebChatAdapter(Platform):
f"[WebChatAdapter] Failed to save proactive message: {e}",
exc_info=True,
)
await super().send_by_session(session, message_chain)
async def _save_proactive_message(
@@ -133,6 +144,7 @@ class WebChatAdapter(Platform):
)
if not message_parts:
return
await db_helper.insert_platform_message_history(
platform_id="webchat",
user_id=conversation_id,
@@ -142,8 +154,7 @@ class WebChatAdapter(Platform):
)
async def _get_message_history(
self,
message_id: int,
self, message_id: int
) -> PlatformMessageHistory | None:
return await db_helper.get_platform_message_history_by_id(message_id)
@@ -153,16 +164,15 @@ class WebChatAdapter(Platform):
depth: int = 0,
max_depth: int = 1,
) -> tuple[list, list[str]]:
"""解析消息段列表,返回消息组件列表和纯文本列表
"""解析消息段列表,返回消息组件列表和纯文本列表
Args:
message_parts: 消息段列表
depth: 当前递归深度
max_depth: 最大递归深度(用于处理 reply)
max_depth: 最大递归深度(用于处理 reply)
Returns:
tuple[list, list[str]]: (消息组件列表, 纯文本列表)
"""
async def get_reply_parts(
@@ -171,10 +181,12 @@ class WebChatAdapter(Platform):
history = await self._get_message_history(message_id)
if not history or not history.content:
return None
reply_parts = history.content.get("message", [])
if not isinstance(reply_parts, list):
return None
return (reply_parts, history.sender_id, history.sender_name)
return reply_parts, history.sender_id, history.sender_name
components, text_parts, _ = await parse_webchat_message_parts(
message_parts,
@@ -186,19 +198,27 @@ class WebChatAdapter(Platform):
max_reply_depth=max_depth,
cast_reply_id_to_str=False,
)
return (components, text_parts)
return components, text_parts
async def convert_message(self, data: tuple) -> AstrBotMessage:
username, cid, payload = data
abm = AstrBotMessage()
abm.self_id = "webchat"
abm.sender = MessageMember(username, username)
abm.type = MessageType.FRIEND_MESSAGE
abm.session_id = f"webchat!{username}!{cid}"
abm.message_id = payload.get("message_id")
# 处理消息段列表
message_parts = payload.get("message", [])
abm.message, message_str_parts = await self._parse_message_parts(message_parts)
logger.debug(f"WebChatAdapter: {abm.message}")
abm.timestamp = int(time.time())
abm.message_str = "".join(message_str_parts)
abm.raw_message = data
@@ -222,18 +242,17 @@ class WebChatAdapter(Platform):
platform_meta=self.meta(),
session_id=message.session_id,
)
_, _, payload = message.raw_message
_, _, payload = message.raw_message # type: ignore
message_event.set_extra("selected_provider", payload.get("selected_provider"))
message_event.set_extra("selected_model", payload.get("selected_model"))
message_event.set_extra(
"enable_streaming",
payload.get("enable_streaming", True),
"enable_streaming", payload.get("enable_streaming", True)
)
message_event.set_extra("action_type", payload.get("action_type"))
message_event.set_extra("llm_checkpoint_id", payload.get("llm_checkpoint_id"))
message_event.set_extra(
"thread_selected_text",
payload.get("thread_selected_text"),
"thread_selected_text", payload.get("thread_selected_text")
)
self.commit_event(message_event)
File diff suppressed because it is too large Load Diff
+154 -252
View File
@@ -7,7 +7,6 @@ import uuid
import wave
from typing import Any
import anyio
import jwt
from quart import websocket
@@ -24,32 +23,12 @@ from astrbot.core.platform.sources.webchat.webchat_queue_mgr import webchat_queu
from astrbot.core.utils.astrbot_path import get_astrbot_data_path, get_astrbot_temp_path
from astrbot.core.utils.datetime_utils import to_utc_isoformat
from .chat import (
BotMessageAccumulator,
build_bot_history_content,
collect_plain_text_from_message_parts,
)
from .route import Route, RouteContext
class LiveChatSession:
"""Live Chat 会话管理器"""
class _ReceiveTimeoutSentinel:
"""Sentinel value indicating a receive timeout."""
_RECEIVE_TIMEOUT = _ReceiveTimeoutSentinel()
class _QueueTimeoutSentinel:
pass
_QUEUE_TIMEOUT = _QueueTimeoutSentinel()
class ClientSession:
def __init__(self, session_id: str, username: str) -> None:
self.session_id = session_id
self.username = username
@@ -77,11 +56,11 @@ class ClientSession:
self.audio_frames.append(data)
async def end_speaking(self, stamp: str) -> tuple[str | None, float]:
"""结束说话,返回组装的 WAV 文件路径和耗时"""
"""结束说话,返回组装的 WAV 文件路径和耗时"""
start_time = time.time()
if not self.is_speaking or stamp != self.current_stamp:
logger.warning(
f"[Live Chat] stamp 不匹配或未在说话状态: {stamp} vs {self.current_stamp}",
f"[Live Chat] stamp 不匹配或未在说话状态: {stamp} vs {self.current_stamp}"
)
return None, 0.0
@@ -97,7 +76,7 @@ class ClientSession:
os.makedirs(temp_dir, exist_ok=True)
audio_path = os.path.join(temp_dir, f"live_audio_{uuid.uuid4()}.wav")
# 假设前端发送的是 PCM 数据,采样率 16000Hz,单声道,16位
# 假设前端发送的是 PCM 数据,采样率 16000Hz,单声道,16位
with wave.open(audio_path, "wb") as wav_file:
wav_file.setnchannels(1) # 单声道
wav_file.setsampwidth(2) # 16位 = 2字节
@@ -107,7 +86,7 @@ class ClientSession:
self.temp_audio_path = audio_path
logger.info(
f"[Live Chat] 音频文件已保存: {audio_path}, 大小: {(await anyio.Path(audio_path).stat()).st_size} bytes",
f"[Live Chat] 音频文件已保存: {audio_path}, 大小: {os.path.getsize(audio_path)} bytes"
)
return audio_path, time.time() - start_time
@@ -115,11 +94,11 @@ class ClientSession:
logger.error(f"[Live Chat] 组装 WAV 文件失败: {e}", exc_info=True)
return None, 0.0
async def cleanup(self) -> None:
def cleanup(self) -> None:
"""清理临时文件"""
if self.temp_audio_path and await anyio.Path(self.temp_audio_path).exists():
if self.temp_audio_path and os.path.exists(self.temp_audio_path):
try:
await anyio.Path(self.temp_audio_path).unlink()
os.remove(self.temp_audio_path)
logger.debug(f"[Live Chat] 已删除临时文件: {self.temp_audio_path}")
except Exception as e:
logger.warning(f"[Live Chat] 删除临时文件失败: {e}")
@@ -140,7 +119,6 @@ class LiveChatRoute(Route):
self.db = db
self.plugin_manager = core_lifecycle.plugin_manager
self.platform_history_mgr = core_lifecycle.platform_message_history_manager
assert self.platform_history_mgr
self.sessions: dict[str, LiveChatSession] = {}
self.attachments_dir = os.path.join(get_astrbot_data_path(), "attachments")
self.legacy_img_dir = os.path.join(get_astrbot_data_path(), "webchat", "imgs")
@@ -151,60 +129,17 @@ class LiveChatRoute(Route):
self.app.websocket("/api/unified_chat/ws")(self.unified_chat_ws)
async def live_chat_ws(self) -> None:
"""Legacy Live Chat WebSocket 处理器(默认 ct=live)"""
"""Legacy Live Chat WebSocket 处理器(默认 ct=live)"""
await self._unified_ws_loop(force_ct="live")
async def unified_chat_ws(self) -> None:
"""Unified Chat WebSocket 处理器(支持 ct=live/chat)"""
"""Unified Chat WebSocket 处理器(支持 ct=live/chat)"""
await self._unified_ws_loop(force_ct=None)
async def _ensure_runtime_ready(self) -> bool:
if is_runtime_request_ready(self.core_lifecycle):
return True
await websocket.close(
1013,
get_runtime_guard_message(self.core_lifecycle),
)
return False
async def _recv_ws_json_guarded(
self,
*,
wait_timeout: float = 1.0,
) -> dict[str, Any] | _ReceiveTimeoutSentinel | None:
if not await self._ensure_runtime_ready():
return None
try:
message = await asyncio.wait_for(
websocket.receive_json(),
timeout=wait_timeout,
)
except asyncio.TimeoutError:
return _RECEIVE_TIMEOUT
if not await self._ensure_runtime_ready():
return None
return message
async def _guarded_queue_get(
self,
back_queue: asyncio.Queue,
*,
wait_timeout: float,
) -> dict[str, Any] | _QueueTimeoutSentinel | None:
if not await self._ensure_runtime_ready():
return None
try:
result = await asyncio.wait_for(back_queue.get(), timeout=wait_timeout)
except asyncio.TimeoutError:
return _QUEUE_TIMEOUT
if not await self._ensure_runtime_ready():
return None
return result
async def _unified_ws_loop(self, force_ct: str | None = None) -> None:
"""统一 WebSocket 循环"""
# WebSocket 不能通过 header 传递 token,需要从 query 参数获取
# 注意:WebSocket 上下文使用 websocket.args 而不是 request.args
# WebSocket 不能通过 header 传递 token,需要从 query 参数获取
# 注意:WebSocket 上下文使用 websocket.args 而不是 request.args
token = websocket.args.get("token")
if not token:
await websocket.close(1008, "Missing authentication token")
@@ -221,9 +156,6 @@ class LiveChatRoute(Route):
await websocket.close(1008, "Invalid token")
return
if not await self._ensure_runtime_ready():
return
session_id = f"webchat_live!{username}!{uuid.uuid4()}"
live_session = LiveChatSession(session_id, username)
self.sessions[session_id] = live_session
@@ -232,11 +164,7 @@ class LiveChatRoute(Route):
try:
while True:
message = await self._recv_ws_json_guarded()
if isinstance(message, _ReceiveTimeoutSentinel):
continue
if message is None:
return
message = await websocket.receive_json()
ct = force_ct or message.get("ct", "live")
if ct == "chat":
await self._handle_chat_message(live_session, message)
@@ -250,16 +178,14 @@ class LiveChatRoute(Route):
# 清理会话
if session_id in self.sessions:
await self._cleanup_chat_subscriptions(live_session)
await live_session.cleanup()
live_session.cleanup()
del self.sessions[session_id]
logger.info(f"[Live Chat] WebSocket 连接关闭: {username}")
async def _create_attachment_from_file(
self,
filename: str,
attach_type: str,
self, filename: str, attach_type: str
) -> dict | None:
"""从本地文件创建 attachment 并返回消息部分。"""
"""从本地文件创建 attachment 并返回消息部分。"""
return await create_attachment_part_from_existing_file(
filename,
attach_type=attach_type,
@@ -269,12 +195,15 @@ class LiveChatRoute(Route):
)
def _extract_web_search_refs(
self,
accumulated_text: str,
accumulated_parts: list,
self, accumulated_text: str, accumulated_parts: list
) -> dict:
"""从消息中提取 web_search 引用。"""
supported = ["web_search_tavily", "web_search_bocha"]
"""从消息中提取 web_search 引用。"""
supported = [
"web_search_baidu",
"web_search_tavily",
"web_search_bocha",
"web_search_brave",
]
web_search_results = {}
tool_call_parts = [
p
@@ -285,7 +214,7 @@ class LiveChatRoute(Route):
for part in tool_call_parts:
for tool_call in part["tool_calls"]:
if tool_call.get("name") not in supported or not tool_call.get(
"result",
"result"
):
continue
try:
@@ -321,21 +250,28 @@ class LiveChatRoute(Route):
async def _save_bot_message(
self,
webchat_conv_id: str,
message_parts: list[dict],
text: str,
media_parts: list,
reasoning: str,
agent_stats: dict,
refs: dict,
llm_checkpoint_id: str | None = None,
):
"""保存 bot 消息到历史记录。"""
new_his = build_bot_history_content(
message_parts,
agent_stats=agent_stats,
refs=refs,
)
bot_message_parts = []
bot_message_parts.extend(media_parts)
if text:
bot_message_parts.append({"type": "plain", "text": text})
mgr = self.platform_history_mgr
assert mgr is not None
return await mgr.insert(
new_his = {"type": "bot", "message": bot_message_parts}
if reasoning:
new_his["reasoning"] = reasoning
if agent_stats:
new_his["agent_stats"] = agent_stats
if refs:
new_his["refs"] = refs
return await self.platform_history_mgr.insert(
platform_id="webchat",
user_id=webchat_conv_id,
content=new_his,
@@ -355,16 +291,11 @@ class LiveChatRoute(Route):
request_id: str,
) -> None:
back_queue = webchat_queue_mgr.get_or_create_back_queue(
request_id,
chat_session_id,
request_id, chat_session_id
)
try:
while True:
result = await self._guarded_queue_get(back_queue, wait_timeout=1)
if isinstance(result, _QueueTimeoutSentinel):
continue
if result is None:
break
result = await back_queue.get()
if not result:
continue
await self._send_chat_payload(session, {"ct": "chat", **result})
@@ -413,11 +344,9 @@ class LiveChatRoute(Route):
session.chat_subscription_tasks.clear()
async def _handle_chat_message(
self,
session: LiveChatSession,
message: dict,
self, session: LiveChatSession, message: dict
) -> None:
"""处理 Chat Mode 消息(ct=chat)"""
"""处理 Chat Mode 消息(ct=chat)"""
msg_type = message.get("t")
if msg_type == "bind":
@@ -528,7 +457,6 @@ class LiveChatRoute(Route):
llm_checkpoint_id = str(uuid.uuid4())
try:
pending_bot_message_flusher = None
chat_queue = webchat_queue_mgr.get_or_create_queue(session_id)
await chat_queue.put(
(
@@ -571,76 +499,22 @@ class LiveChatRoute(Route):
},
)
message_accumulator = BotMessageAccumulator()
accumulated_parts = []
accumulated_text = ""
accumulated_reasoning = ""
tool_calls = {}
agent_stats = {}
refs: dict[str, Any] = {}
async def flush_pending_bot_message():
nonlocal message_accumulator, agent_stats, refs
if not (message_accumulator.has_content() or refs or agent_stats):
return None
message_parts_to_save = message_accumulator.build_message_parts(
include_pending_tool_calls=True
)
plain_text = collect_plain_text_from_message_parts(
message_parts_to_save
)
try:
extracted_refs = self._extract_web_search_refs(
plain_text,
message_parts_to_save,
)
except Exception as e:
logger.exception(
f"[Live Chat] Failed to extract web search refs: {e}",
exc_info=True,
)
extracted_refs = refs
saved_record = await self._save_bot_message(
session_id,
message_parts_to_save,
agent_stats,
extracted_refs,
llm_checkpoint_id,
)
message_accumulator = BotMessageAccumulator()
agent_stats = {}
refs = {}
return saved_record
pending_bot_message_flusher = flush_pending_bot_message
async def send_attachment_saved_event(part: dict | None) -> None:
if not part or not part.get("attachment_id") or not part.get("type"):
return
await self._send_chat_payload(
session,
{
"ct": "chat",
"type": "attachment_saved",
"data": {
"id": part["attachment_id"],
"type": part["type"],
},
},
)
refs = {}
while True:
if not await self._ensure_runtime_ready():
break
if session.should_interrupt:
session.should_interrupt = False
await flush_pending_bot_message()
break
result = await self._guarded_queue_get(back_queue, wait_timeout=1)
if isinstance(result, _QueueTimeoutSentinel):
try:
result = await asyncio.wait_for(back_queue.get(), timeout=1)
except asyncio.TimeoutError:
continue
if result is None:
break
if not result:
continue
@@ -671,36 +545,68 @@ class LiveChatRoute(Route):
await self._send_chat_payload(session, outgoing)
if msg_type == "plain":
message_accumulator.add_plain(
result_text,
chain_type=chain_type,
streaming=streaming,
)
if chain_type == "tool_call":
try:
tool_call = json.loads(result_text)
tool_calls[tool_call.get("id")] = tool_call
if accumulated_text:
accumulated_parts.append(
{"type": "plain", "text": accumulated_text}
)
accumulated_text = ""
except Exception:
pass
elif chain_type == "tool_call_result":
try:
tcr = json.loads(result_text)
tc_id = tcr.get("id")
if tc_id in tool_calls:
tool_calls[tc_id]["result"] = tcr.get("result")
tool_calls[tc_id]["finished_ts"] = tcr.get("ts")
accumulated_parts.append(
{
"type": "tool_call",
"tool_calls": [tool_calls[tc_id]],
}
)
tool_calls.pop(tc_id, None)
except Exception:
pass
elif chain_type == "reasoning":
accumulated_reasoning += result_text
elif streaming:
accumulated_text += result_text
else:
accumulated_text = result_text
elif msg_type == "image":
filename = str(result_text).replace("[IMAGE]", "")
part = await self._create_attachment_from_file(filename, "image")
message_accumulator.add_attachment(part)
await send_attachment_saved_event(part)
if part:
accumulated_parts.append(part)
elif msg_type == "record":
filename = str(result_text).replace("[RECORD]", "")
part = await self._create_attachment_from_file(filename, "record")
message_accumulator.add_attachment(part)
await send_attachment_saved_event(part)
if part:
accumulated_parts.append(part)
elif msg_type == "file":
filename = str(result_text).replace("[FILE]", "").split("|", 1)[0]
part = await self._create_attachment_from_file(filename, "file")
message_accumulator.add_attachment(part)
await send_attachment_saved_event(part)
if part:
accumulated_parts.append(part)
elif msg_type == "video":
filename = str(result_text).replace("[VIDEO]", "").split("|", 1)[0]
part = await self._create_attachment_from_file(filename, "video")
message_accumulator.add_attachment(part)
await send_attachment_saved_event(part)
if part:
accumulated_parts.append(part)
should_save = False
if msg_type == "end":
should_save = bool(
message_accumulator.has_content() or refs or agent_stats
accumulated_parts
or accumulated_text
or accumulated_reasoning
or refs
or agent_stats
)
elif (streaming and msg_type == "complete") or not streaming:
if chain_type not in (
@@ -711,7 +617,26 @@ class LiveChatRoute(Route):
should_save = True
if should_save:
saved_record = await flush_pending_bot_message()
try:
refs = self._extract_web_search_refs(
accumulated_text,
accumulated_parts,
)
except Exception as e:
logger.exception(
f"[Live Chat] Failed to extract web search refs: {e}",
exc_info=True,
)
saved_record = await self._save_bot_message(
session_id,
accumulated_text,
accumulated_parts,
accumulated_reasoning,
agent_stats,
refs,
llm_checkpoint_id,
)
if saved_record:
await self._send_chat_payload(
session,
@@ -721,13 +646,19 @@ class LiveChatRoute(Route):
"data": {
"id": saved_record.id,
"created_at": to_utc_isoformat(
saved_record.created_at,
saved_record.created_at
),
"llm_checkpoint_id": llm_checkpoint_id,
},
},
)
accumulated_parts = []
accumulated_text = ""
accumulated_reasoning = ""
agent_stats = {}
refs = {}
if msg_type == "end":
break
@@ -738,24 +669,16 @@ class LiveChatRoute(Route):
{
"ct": "chat",
"t": "error",
"data": f"处理失败: {e!s}",
"data": f"处理失败: {str(e)}",
"code": "PROCESSING_ERROR",
},
)
finally:
try:
if pending_bot_message_flusher is not None:
await pending_bot_message_flusher()
except Exception as e:
logger.exception(
f"[Live Chat] Failed to persist pending chat message: {e}",
exc_info=True,
)
session.is_processing = False
webchat_queue_mgr.remove_back_queue(message_id)
async def _build_chat_message_parts(self, message: list[dict]) -> list[dict]:
"""构建 chat websocket 用户消息段(复用 webchat 逻辑)"""
"""构建 chat websocket 用户消息段(复用 webchat 逻辑)"""
return await build_webchat_message_parts(
message,
get_attachment_by_id=self.db.get_attachment_by_id,
@@ -801,7 +724,7 @@ class LiveChatRoute(Route):
await websocket.send_json({"t": "error", "data": "音频组装失败"})
return
# 处理音频:STT -> LLM -> TTS
# 处理音频:STT -> LLM -> TTS
await self._process_audio(session, audio_path, assemble_duration)
elif msg_type == "interrupt":
@@ -810,16 +733,13 @@ class LiveChatRoute(Route):
logger.info(f"[Live Chat] 用户打断: {session.username}")
async def _process_audio(
self,
session: LiveChatSession,
audio_path: str,
assemble_duration: float,
self, session: LiveChatSession, audio_path: str, assemble_duration: float
) -> None:
"""处理音频:STT -> LLM -> 流式 TTS"""
"""处理音频:STT -> LLM -> 流式 TTS"""
try:
# 发送 WAV 组装耗时
await websocket.send_json(
{"t": "metrics", "data": {"wav_assemble_time": assemble_duration}},
{"t": "metrics", "data": {"wav_assemble_time": assemble_duration}}
)
wav_assembly_finish_time = time.time()
@@ -827,14 +747,7 @@ class LiveChatRoute(Route):
session.should_interrupt = False
# 1. STT - 语音转文字
pm = self.plugin_manager
if pm is None or pm.context is None:
logger.error("[Live Chat] Plugin manager not available")
await websocket.send_json(
{"t": "error", "data": "Plugin manager not available"},
)
return
ctx = pm.context
ctx = self.plugin_manager.context
stt_provider = ctx.provider_manager.stt_provider_insts[0]
if not stt_provider:
@@ -843,7 +756,7 @@ class LiveChatRoute(Route):
return
await websocket.send_json(
{"t": "metrics", "data": {"stt": stt_provider.meta().type}},
{"t": "metrics", "data": {"stt": stt_provider.meta().type}}
)
user_text = await stt_provider.get_text(audio_path)
@@ -857,7 +770,7 @@ class LiveChatRoute(Route):
{
"t": "user_msg",
"data": {"text": user_text, "ts": int(time.time() * 1000)},
},
}
)
# 2. 构造消息事件并发送到 pipeline
@@ -883,17 +796,13 @@ class LiveChatRoute(Route):
try:
while True:
if not await self._ensure_runtime_ready():
break
if session.should_interrupt:
# 用户打断,停止处理
# 用户打断,停止处理
logger.info("[Live Chat] 检测到用户打断")
await websocket.send_json({"t": "stop_play"})
# 保存消息并标记为被打断
await self._save_interrupted_message(
session,
user_text,
bot_text,
session, user_text, bot_text
)
# 清空队列中未处理的消息
while not back_queue.empty():
@@ -903,14 +812,10 @@ class LiveChatRoute(Route):
break
break
result = await self._guarded_queue_get(
back_queue,
wait_timeout=0.5,
)
if isinstance(result, _QueueTimeoutSentinel):
try:
result = await asyncio.wait_for(back_queue.get(), timeout=0.5)
except asyncio.TimeoutError:
continue
if result is None:
break
if not result:
continue
@@ -918,7 +823,7 @@ class LiveChatRoute(Route):
result_message_id = result.get("message_id")
if result_message_id != message_id:
logger.warning(
f"[Live Chat] 消息 ID 不匹配: {result_message_id} != {message_id}",
f"[Live Chat] 消息 ID 不匹配: {result_message_id} != {message_id}"
)
continue
@@ -937,7 +842,7 @@ class LiveChatRoute(Route):
"llm_total_time": stats.get("end_time", 0)
- stats.get("start_time", 0),
},
},
}
)
except Exception as e:
logger.error(f"[Live Chat] 解析 AgentStats 失败: {e}")
@@ -950,7 +855,7 @@ class LiveChatRoute(Route):
{
"t": "metrics",
"data": stats,
},
}
)
except Exception as e:
logger.error(f"[Live Chat] 解析 TTSStats 失败: {e}")
@@ -974,9 +879,9 @@ class LiveChatRoute(Route):
{
"t": "metrics",
"data": {
"speak_to_first_frame": speak_to_first_frame_latency,
"speak_to_first_frame": speak_to_first_frame_latency
},
},
}
)
text = result.get("text")
@@ -985,7 +890,7 @@ class LiveChatRoute(Route):
{
"t": "bot_text_chunk",
"data": {"text": text},
},
}
)
# 发送音频数据给前端
@@ -993,14 +898,14 @@ class LiveChatRoute(Route):
{
"t": "response",
"data": data, # base64 编码的音频数据
},
}
)
elif result_type in ["complete", "end"]:
# 处理完成
logger.info(f"[Live Chat] Bot 回复完成: {bot_text}")
# 如果没有音频流,发送 bot 消息文本
# 如果没有音频流,发送 bot 消息文本
if not audio_playing:
await websocket.send_json(
{
@@ -1009,7 +914,7 @@ class LiveChatRoute(Route):
"text": bot_text,
"ts": int(time.time() * 1000),
},
},
}
)
# 发送结束标记
@@ -1021,7 +926,7 @@ class LiveChatRoute(Route):
{
"t": "metrics",
"data": {"wav_to_tts_total_time": wav_to_tts_duration},
},
}
)
break
finally:
@@ -1029,31 +934,28 @@ class LiveChatRoute(Route):
except Exception as e:
logger.error(f"[Live Chat] 处理音频失败: {e}", exc_info=True)
await websocket.send_json({"t": "error", "data": f"处理失败: {e!s}"})
await websocket.send_json({"t": "error", "data": f"处理失败: {str(e)}"})
finally:
session.is_processing = False
session.should_interrupt = False
async def _save_interrupted_message(
self,
session: LiveChatSession,
user_text: str,
bot_text: str,
self, session: LiveChatSession, user_text: str, bot_text: str
) -> None:
"""保存被打断的消息"""
interrupted_text = bot_text + " [用户打断]"
logger.info(f"[Live Chat] 保存打断消息: {interrupted_text}")
# 简单记录到日志,实际保存逻辑可以后续完善
# 简单记录到日志,实际保存逻辑可以后续完善
try:
timestamp = int(time.time() * 1000)
logger.info(
f"[Live Chat] 用户消息: {user_text} (session: {session.session_id}, ts: {timestamp})",
f"[Live Chat] 用户消息: {user_text} (session: {session.session_id}, ts: {timestamp})"
)
if bot_text:
logger.info(
f"[Live Chat] Bot 消息(打断): {interrupted_text} (session: {session.session_id}, ts: {timestamp})",
f"[Live Chat] Bot 消息(打断): {interrupted_text} (session: {session.session_id}, ts: {timestamp})"
)
except Exception as e:
logger.error(f"[Live Chat] 记录消息失败: {e}", exc_info=True)
File diff suppressed because it is too large Load Diff
+141 -149
View File
@@ -1,6 +1,7 @@
<template>
<div
class="input-area fade-in"
:class="{ 'is-dark': isDark }"
@dragover.prevent="handleDragOver"
@dragleave.prevent="handleDragLeave"
@drop.prevent="handleDrop"
@@ -14,45 +15,41 @@
border: isDark ? 'none' : '1px solid #e0e0e0',
borderRadius: '24px',
boxShadow: isDark ? 'none' : '0px 2px 2px rgba(0, 0, 0, 0.1)',
backgroundColor: isDark ? 'rgba(15, 15, 22, 0.6)' : 'transparent',
backgroundColor: isDark ? '#2d2d2d' : 'transparent',
position: 'relative',
transition: 'min-height 0.2s ease, padding 0.2s ease',
}"
>
<!-- 拖拽上传遮罩 -->
<transition name="fade">
<div v-if="isDragging" class="drop-overlay">
<div class="drop-overlay-content">
<v-icon size="48" color="primary"> mdi-cloud-upload </v-icon>
<v-icon size="48" color="primary">mdi-cloud-upload</v-icon>
<span class="drop-text">{{ tm("input.dropToUpload") }}</span>
</div>
</div>
</transition>
<!-- 引用预览区 -->
<transition name="slideReply" @after-leave="handleReplyAfterLeave">
<div v-if="props.replyTo && !isReplyClosing" class="reply-preview">
<div class="reply-preview" v-if="props.replyTo && !isReplyClosing">
<div class="reply-content">
<v-icon size="small" class="reply-icon"> mdi-reply </v-icon>
<v-icon size="small" class="reply-icon">mdi-reply</v-icon>
"<span class="reply-text">{{ props.replyTo.selectedText }}</span
>"
</div>
<v-btn
@click="handleClearReply"
class="remove-reply-btn"
icon="mdi-close"
size="x-small"
color="grey"
variant="text"
@click="handleClearReply"
/>
</div>
</transition>
<textarea
ref="inputField"
v-model="localPrompt"
@compositionstart="handleCompositionStart"
@compositionend="handleCompositionEnd"
@compositioncancel="handleCompositionEnd"
@blur="clearCompositionState()"
@keydown="handleKeyDown"
:disabled="disabled"
placeholder="Ask AstrBot..."
class="chat-textarea"
@@ -66,17 +63,15 @@
outline: none;
border: 1px solid var(--v-theme-border);
border-radius: 12px;
padding: 16px 20px;
min-height: 40px;
padding: 12px 18px;
min-height: 34px;
max-height: 200px;
overflow-y: auto;
font-family: inherit;
font-size: 16px;
background-color: var(--v-theme-surface);
transition: height 0.16s ease;
"
@keydown="handleKeyDown"
/>
></textarea>
<div
style="
display: flex;
@@ -103,12 +98,12 @@
location="top start"
:close-on-content-click="false"
>
<template #activator="{ props: activatorProps }">
<template v-slot:activator="{ props: activatorProps }">
<v-btn
v-bind="activatorProps"
icon="mdi-plus"
variant="text"
color="primary"
variant="outlined"
class="input-neutral-btn input-outline-control"
/>
</template>
@@ -118,8 +113,8 @@
rounded="md"
@click="triggerImageInput"
>
<template #prepend>
<v-icon icon="mdi-file-upload-outline" size="small" />
<template v-slot:prepend>
<v-icon icon="mdi-file-upload" size="small"></v-icon>
</template>
<v-list-item-title>
{{ tm("input.upload") }}
@@ -141,11 +136,8 @@
rounded="md"
@click="$emit('toggleStreaming')"
>
<template #prepend>
<v-icon
:icon="enableStreaming ? 'mdi-flash' : 'mdi-flash-off'"
size="small"
/>
<template v-slot:prepend>
<v-icon icon="mdi-lightning-bolt" size="small"></v-icon>
</template>
<v-list-item-title>
{{
@@ -173,11 +165,11 @@
"
>
<input
ref="imageInputRef"
type="file"
ref="imageInputRef"
@change="handleFileSelect"
style="display: none"
multiple
@change="handleFileSelect"
/>
<v-progress-circular
v-if="disabled && !mobile"
@@ -189,7 +181,7 @@
<!-- <v-btn @click="$emit('openLiveMode')"
icon
variant="text"
color="purple"
color="purple"
size="small"
>
<v-icon icon="mdi-phone-in-talk" variant="text" plain></v-icon>
@@ -198,17 +190,16 @@
</v-tooltip>
</v-btn> -->
<v-btn
@click="handleRecordClick"
icon
variant="text"
:color="isRecording ? 'error' : 'primary'"
class="record-btn"
@click="handleRecordClick"
class="record-btn input-icon-btn"
>
<v-icon
:icon="isRecording ? 'mdi-stop-circle' : 'mdi-microphone'"
variant="text"
plain
/>
></v-icon>
<v-tooltip activator="parent" location="top">
{{
isRecording ? tm("voice.speaking") : tm("voice.startRecording")
@@ -216,26 +207,24 @@
</v-tooltip>
</v-btn>
<v-btn
v-if="isRunning && !canSend"
icon
variant="tonal"
color="primary"
class="send-btn"
v-if="isRunning && !canSend"
@click="$emit('stop')"
variant="tonal"
class="send-btn input-action-btn"
>
<v-icon icon="mdi-stop" variant="text" plain />
<v-icon icon="mdi-stop" variant="text" plain></v-icon>
<v-tooltip activator="parent" location="top">
{{ tm("input.stopGenerating") }}
</v-tooltip>
</v-btn>
<v-btn
v-else
icon="mdi-send"
variant="tonal"
color="primary"
:disabled="!canSend"
class="send-btn"
@click="$emit('send')"
icon="mdi-arrow-up"
variant="tonal"
:disabled="!canSend"
class="send-btn input-action-btn"
/>
</div>
</div>
@@ -243,12 +232,12 @@
<!-- 附件预览区 -->
<div
class="attachments-preview"
v-if="
stagedImagesUrl.length > 0 ||
stagedAudioUrl ||
(stagedFiles && stagedFiles.length > 0)
"
class="attachments-preview"
>
<div
v-for="(img, index) in stagedImagesUrl"
@@ -257,27 +246,27 @@
>
<img :src="img" class="preview-image" />
<v-btn
@click="$emit('removeImage', index)"
class="remove-attachment-btn"
icon="mdi-close"
size="small"
color="error"
variant="text"
@click="$emit('removeImage', index)"
/>
</div>
<div v-if="stagedAudioUrl" class="audio-preview">
<v-chip color="primary" variant="tonal" class="audio-chip">
<v-icon start icon="mdi-microphone" size="small" />
<v-icon start icon="mdi-microphone" size="small"></v-icon>
{{ tm("voice.recording") }}
</v-chip>
<v-btn
@click="$emit('removeAudio')"
class="remove-attachment-btn"
icon="mdi-close"
size="small"
color="error"
variant="text"
@click="$emit('removeAudio')"
/>
</div>
@@ -287,16 +276,16 @@
class="file-preview"
>
<v-chip color="primary" variant="tonal" class="file-chip">
<v-icon start icon="mdi-file-document-outline" size="small" />
<v-icon start icon="mdi-file-document-outline" size="small"></v-icon>
<span class="file-name-preview">{{ file.original_name }}</span>
</v-chip>
<v-btn
@click="$emit('removeFile', index)"
class="remove-attachment-btn"
icon="mdi-close"
size="small"
color="error"
variant="text"
@click="$emit('removeFile', index)"
/>
</div>
</div>
@@ -315,7 +304,6 @@ import {
import { useDisplay } from "vuetify";
import { useModuleI18n } from "@/i18n/composables";
import { useCustomizerStore } from "@/stores/customizer";
import { isComposingEnter } from "@/utils/imeInput.mjs";
import ConfigSelector from "./ConfigSelector.vue";
import ProviderModelMenu from "./ProviderModelMenu.vue";
import StyledMenu from "@/components/shared/StyledMenu.vue";
@@ -330,7 +318,7 @@ interface StagedFileInfo {
}
interface ReplyInfo {
messageId: number;
messageId: string | number;
selectedText?: string;
}
@@ -376,8 +364,9 @@ const emit = defineEmits<{
}>();
const { tm } = useModuleI18n("features/chat");
// 从新的预设getter获取
const isDark = computed(() => useCustomizerStore().isDarkTheme);
const isDark = computed(
() => useCustomizerStore().uiTheme === "PurpleThemeDark",
);
const inputField = ref<HTMLTextAreaElement | null>(null);
const imageInputRef = ref<HTMLInputElement | null>(null);
@@ -387,8 +376,6 @@ const providerModelMenuRef = ref<InstanceType<typeof ProviderModelMenu> | null>(
const showProviderSelector = ref(true);
const isReplyClosing = ref(false);
const isDragging = ref(false);
const isComposing = ref(false);
const lastCompositionEndAt = ref<number | null>(null);
let dragLeaveTimeout: number | null = null;
const localPrompt = computed({
@@ -410,44 +397,6 @@ const canSend = computed(() => {
);
});
const fileTypeStyles: Record<
string,
{ color: string; icon: string; label: string }
> = {
pdf: { color: "#d32f2f", icon: "mdi-file-pdf-box", label: "PDF" },
txt: { color: "#1976d2", icon: "mdi-file-document-outline", label: "TXT" },
md: { color: "#1976d2", icon: "mdi-language-markdown-outline", label: "MD" },
doc: { color: "#2b579a", icon: "mdi-file-word-box", label: "DOC" },
docx: { color: "#2b579a", icon: "mdi-file-word-box", label: "DOCX" },
xls: { color: "#217346", icon: "mdi-file-excel-box", label: "XLS" },
xlsx: { color: "#217346", icon: "mdi-file-excel-box", label: "XLSX" },
csv: { color: "#217346", icon: "mdi-file-delimited-outline", label: "CSV" },
zip: { color: "#7b5e00", icon: "mdi-folder-zip-outline", label: "ZIP" },
py: { color: "#3776ab", icon: "mdi-language-python", label: "PY" },
js: { color: "#b8860b", icon: "mdi-language-javascript", label: "JS" },
ts: { color: "#3178c6", icon: "mdi-language-typescript", label: "TS" },
html: { color: "#e34c26", icon: "mdi-language-html5", label: "HTML" },
css: { color: "#264de4", icon: "mdi-language-css3", label: "CSS" },
json: { color: "#6a1b9a", icon: "mdi-code-json", label: "JSON" },
};
function fileExtension(file: StagedFileInfo) {
const name = file.original_name || file.filename || "";
const extension = name.split(".").pop()?.toLowerCase() || "";
return extension === name.toLowerCase() ? "" : extension;
}
function filePresentation(file: StagedFileInfo) {
const extension = fileExtension(file);
return (
fileTypeStyles[extension] || {
color: "#607d8b",
icon: "mdi-file-document-outline",
label: extension ? extension.slice(0, 4).toUpperCase() : "FILE",
}
);
}
// Ctrl+B 长按录音相关
const ctrlKeyDown = ref(false);
const ctrlKeyTimer = ref<number | null>(null);
@@ -496,10 +445,6 @@ function handleKeyDown(e: KeyboardEvent) {
return;
}
if (isComposingEnter(e, isComposing.value, lastCompositionEndAt.value)) {
return;
}
const isSendHotkey =
e.ctrlKey ||
e.metaKey ||
@@ -519,23 +464,6 @@ function handleKeyDown(e: KeyboardEvent) {
}
}
function handleCompositionStart() {
isComposing.value = true;
lastCompositionEndAt.value = null;
}
function handleCompositionEnd(e: CompositionEvent) {
lastCompositionEndAt.value = e.timeStamp;
clearCompositionState({ keepLastEndAt: true });
}
function clearCompositionState({ keepLastEndAt = false } = {}) {
isComposing.value = false;
if (!keepLastEndAt) {
lastCompositionEndAt.value = null;
}
}
function handleKeyUp(e: KeyboardEvent) {
if (e.keyCode === 66) {
ctrlKeyDown.value = false;
@@ -637,7 +565,6 @@ onBeforeUnmount(() => {
if (inputField.value) {
inputField.value.removeEventListener("paste", handlePaste);
}
clearCompositionState();
document.removeEventListener("keyup", handleKeyUp);
});
@@ -648,44 +575,101 @@ defineExpose({
</script>
<style scoped>
/* Dark mode input container glass */
.v-theme--bluebusinessdarktheme :deep(.input-container),
.v-theme--bluebusinessdarktheme ::v-deep(.input-container) {
background: rgba(15, 15, 22, 0.6) !important;
backdrop-filter: blur(16px) saturate(1.2) !important;
border: 1px solid rgba(0, 242, 255, 0.08) !important;
box-shadow: 0 0 20px rgba(0, 0, 0, 0.3) !important;
}
/* Light mode: clean white frosted glass */
.v-theme--bluebusinesstheme :deep(.input-container),
.v-theme--bluebusinesstheme ::v-deep(.input-container) {
background: rgba(255, 255, 255, 0.9) !important;
backdrop-filter: blur(20px) saturate(1.1) !important;
border: 1px solid rgba(0, 49, 83, 0.1) !important;
box-shadow: 0 2px 12px rgba(26, 46, 60, 0.08) !important;
}
/* Fix placeholder visibility in dark mode */
.v-theme--bluebusinessdarktheme .chat-textarea::placeholder {
color: rgba(228, 225, 230, 0.35) !important;
opacity: 1 !important;
}
/* Fix placeholder visibility in light mode */
.v-theme--bluebusinesstheme .chat-textarea::placeholder {
color: rgba(26, 46, 80, 0.35) !important;
opacity: 1 !important;
}
.input-area {
padding: 16px;
padding: 12px 16px 0;
background-color: transparent;
position: relative;
border-top: 1px solid var(--v-theme-border);
flex-shrink: 0;
}
.input-neutral-btn {
color: #6f6f6f !important;
}
.input-neutral-btn:hover {
background: #efefef;
}
.input-neutral-btn--tonal {
background: #efefef;
color: #4f4f4f !important;
}
.input-neutral-btn--tonal:hover {
background: #e7e7e7;
}
.input-action-btn {
background: #5594c6 !important;
color: #fff !important;
}
.input-action-btn:hover {
background: #4c86b3 !important;
}
.input-action-btn:disabled {
background: rgba(85, 148, 198, 0.24) !important;
color: rgba(255, 255, 255, 0.72) !important;
}
.input-icon-btn {
background: transparent !important;
color: rgb(var(--v-theme-on-surface)) !important;
margin-right: 8px;
}
.input-icon-btn:hover {
background: rgba(var(--v-theme-on-surface), 0.04) !important;
}
.input-outline-control {
width: 36px !important;
height: 36px !important;
min-width: 36px !important;
border-color: rgba(var(--v-theme-on-surface), 0.18) !important;
background: transparent !important;
}
.input-outline-control:hover {
border-color: rgba(var(--v-theme-on-surface), 0.34) !important;
background: rgba(var(--v-theme-on-surface), 0.04) !important;
}
.input-area.is-dark .input-neutral-btn {
color: rgba(255, 255, 255, 0.78) !important;
}
.input-area.is-dark .input-neutral-btn:hover,
.input-area.is-dark .input-neutral-btn--tonal {
background: rgba(255, 255, 255, 0.1);
}
.input-area.is-dark .input-outline-control {
border-color: rgba(255, 255, 255, 0.22) !important;
background: transparent !important;
}
.input-area.is-dark .input-outline-control:hover {
border-color: rgba(255, 255, 255, 0.42) !important;
background: rgba(255, 255, 255, 0.06) !important;
}
.input-area.is-dark .input-action-btn {
background: rgb(var(--v-theme-on-surface)) !important;
color: rgb(var(--v-theme-surface)) !important;
}
.input-area.is-dark .input-action-btn:hover {
background: rgba(var(--v-theme-on-surface), 0.86) !important;
}
.input-area.is-dark .input-action-btn:disabled {
background: rgba(var(--v-theme-on-surface), 0.14) !important;
color: rgba(var(--v-theme-on-surface), 0.4) !important;
}
/* 拖拽上传遮罩 */
.drop-overlay {
position: absolute;
@@ -882,20 +866,28 @@ defineExpose({
@media (max-width: 768px) {
.input-area {
padding: 0 !important;
padding-bottom: 10px !important;
}
.input-container {
width: 100% !important;
max-width: 100% !important;
border-bottom-left-radius: 0 !important;
border-bottom-right-radius: 0 !important;
}
.input-outline-control {
width: 32px !important;
height: 32px !important;
min-width: 32px !important;
}
.input-area textarea,
.chat-textarea {
min-height: 32px !important;
max-height: 160px !important;
min-height: 28px !important;
max-height: 140px !important;
font-size: 16px !important;
padding: 16px 16px 12px 16px !important;
line-height: 20px !important;
padding: 8px 14px 7px !important;
}
}
</style>
@@ -24,54 +24,6 @@
<div class="message-stack">
<div
v-if="isUserMessage(msg) && userAttachmentParts(msg).length"
class="sent-attachments"
:class="{ 'images-only': hasImageOnlyAttachments(msg) }"
>
<template
v-for="(part, attachmentIndex) in userAttachmentParts(msg)"
:key="`${msgIndex}-attachment-${attachmentIndex}-${part.type}`"
>
<button
v-if="part.type === 'image'"
class="sent-attachment-card sent-image-card"
type="button"
@click="openImage(partUrl(part))"
>
<img :src="partUrl(part)" :alt="part.filename || 'image'" />
</button>
<div v-else class="sent-attachment-card sent-file-card">
<div
class="sent-attachment-icon"
:style="{ color: attachmentPresentation(part).color }"
>
<v-icon :icon="attachmentPresentation(part).icon" size="24" />
<span class="sent-attachment-ext">
{{ attachmentPresentation(part).label }}
</span>
</div>
<span class="sent-attachment-name">
{{ attachmentName(part) }}
</span>
<v-btn
v-if="part.type === 'file'"
icon="mdi-download"
size="x-small"
variant="text"
:loading="
downloadingFiles.has(
part.attachment_id || part.filename || '',
)
"
@click="downloadPart(part)"
/>
</div>
</template>
</div>
<div
v-if="shouldShowMessageBubble(msg)"
class="message-bubble"
:class="{ user: isUserMessage(msg), bot: !isUserMessage(msg) }"
@mouseup="handleMouseUp($event, msg)"
@@ -129,7 +81,7 @@
/>
<template
v-for="(part, partIndex) in bubbleParts(msg)"
v-for="(part, partIndex) in messageParts(msg)"
:key="`${msgIndex}-${partIndex}-${part.type}`"
>
<button
@@ -465,37 +417,6 @@ function messageParts(message: ChatRecord): MessagePart[] {
return [];
}
function isAttachmentPart(part: MessagePart) {
return ["image", "record", "video", "file"].includes(part.type);
}
function userAttachmentParts(message: ChatRecord) {
if (!isUserMessage(message)) return [];
return messageParts(message).filter(isAttachmentPart);
}
function hasImageOnlyAttachments(message: ChatRecord) {
const attachments = userAttachmentParts(message);
return (
attachments.length > 0 &&
attachments.every((part) => part.type === "image")
);
}
function bubbleParts(message: ChatRecord) {
if (!isUserMessage(message)) return messageParts(message);
return messageParts(message).filter((part) => !isAttachmentPart(part));
}
function shouldShowMessageBubble(message: ChatRecord) {
return (
!isUserMessage(message) ||
isEditingMessage(message) ||
messageContent(message).isLoading ||
bubbleParts(message).length > 0
);
}
function isMessageStreaming(message: ChatRecord, messageIndex: number) {
return (
props.isStreaming &&
@@ -561,74 +482,6 @@ function hasNonReasoningContent(message: ChatRecord) {
});
}
const attachmentTypeStyles: Record<
string,
{ color: string; icon: string; label: string }
> = {
pdf: { color: "#d32f2f", icon: "mdi-file-pdf-box", label: "PDF" },
txt: { color: "#1976d2", icon: "mdi-file-document-outline", label: "TXT" },
md: { color: "#1976d2", icon: "mdi-language-markdown-outline", label: "MD" },
markdown: {
color: "#1976d2",
icon: "mdi-language-markdown-outline",
label: "MD",
},
doc: { color: "#2b579a", icon: "mdi-file-word-box", label: "DOC" },
docx: { color: "#2b579a", icon: "mdi-file-word-box", label: "DOCX" },
xls: { color: "#217346", icon: "mdi-file-excel-box", label: "XLS" },
xlsx: { color: "#217346", icon: "mdi-file-excel-box", label: "XLSX" },
csv: { color: "#217346", icon: "mdi-file-delimited-outline", label: "CSV" },
ppt: { color: "#d24726", icon: "mdi-file-powerpoint-box", label: "PPT" },
pptx: { color: "#d24726", icon: "mdi-file-powerpoint-box", label: "PPTX" },
zip: { color: "#7b5e00", icon: "mdi-folder-zip-outline", label: "ZIP" },
rar: { color: "#7b5e00", icon: "mdi-folder-zip-outline", label: "RAR" },
"7z": { color: "#7b5e00", icon: "mdi-folder-zip-outline", label: "7Z" },
tar: { color: "#7b5e00", icon: "mdi-folder-zip-outline", label: "TAR" },
gz: { color: "#7b5e00", icon: "mdi-folder-zip-outline", label: "GZ" },
json: { color: "#6a1b9a", icon: "mdi-code-json", label: "JSON" },
yaml: { color: "#6a1b9a", icon: "mdi-code-braces", label: "YAML" },
yml: { color: "#6a1b9a", icon: "mdi-code-braces", label: "YML" },
js: { color: "#b8860b", icon: "mdi-language-javascript", label: "JS" },
ts: { color: "#3178c6", icon: "mdi-language-typescript", label: "TS" },
html: { color: "#e34c26", icon: "mdi-language-html5", label: "HTML" },
css: { color: "#264de4", icon: "mdi-language-css3", label: "CSS" },
py: { color: "#3776ab", icon: "mdi-language-python", label: "PY" },
java: { color: "#b07219", icon: "mdi-language-java", label: "JAVA" },
mp3: { color: "#00897b", icon: "mdi-file-music-outline", label: "MP3" },
wav: { color: "#00897b", icon: "mdi-file-music-outline", label: "WAV" },
flac: { color: "#00897b", icon: "mdi-file-music-outline", label: "FLAC" },
mp4: { color: "#5e35b1", icon: "mdi-file-video-outline", label: "MP4" },
mov: { color: "#5e35b1", icon: "mdi-file-video-outline", label: "MOV" },
webm: { color: "#5e35b1", icon: "mdi-file-video-outline", label: "WEBM" },
};
function attachmentName(part: MessagePart) {
return part.embedded_file?.filename || part.filename || part.type || "file";
}
function attachmentExtension(part: MessagePart) {
const name = attachmentName(part);
const extension = name.split(".").pop()?.toLowerCase() || "";
return extension === name.toLowerCase() ? "" : extension;
}
function attachmentPresentation(part: MessagePart) {
if (part.type === "record") {
return { color: "#00897b", icon: "mdi-microphone", label: "AUDIO" };
}
if (part.type === "video") {
return { color: "#5e35b1", icon: "mdi-file-video-outline", label: "VIDEO" };
}
const extension = attachmentExtension(part);
return (
attachmentTypeStyles[extension] || {
color: "#607d8b",
icon: "mdi-file-document-outline",
label: extension ? extension.slice(0, 4).toUpperCase() : "FILE",
}
);
}
function handleMouseUp(event: MouseEvent, message: ChatRecord) {
if (props.enableThreadSelection && !isUserMessage(message)) {
emit("selectBotText", event, message);
@@ -895,8 +748,6 @@ function formatDuration(seconds: number) {
}
.message-stack {
display: flex;
flex-direction: column;
max-width: min(760px, 82%);
}
@@ -905,95 +756,6 @@ function formatDuration(seconds: number) {
max-width: 60%;
}
.sent-attachments {
display: flex;
max-width: 100%;
gap: 10px;
margin-bottom: 8px;
padding: 2px 2px 4px;
overflow-x: auto;
overflow-y: hidden;
scrollbar-width: thin;
}
.sent-attachment-card {
position: relative;
display: inline-flex;
flex: 0 0 auto;
align-items: center;
justify-content: flex-start;
gap: 8px;
height: 64px;
overflow: hidden;
border: 1px solid rgba(var(--v-theme-on-surface), 0.1);
border-radius: 12px;
background: rgba(var(--v-theme-on-surface), 0.04);
color: rgb(var(--v-theme-on-surface));
}
.sent-image-card {
width: 64px;
padding: 0;
border: 0;
cursor: zoom-in;
}
.sent-image-card img {
width: 100%;
height: 100%;
border-radius: 11px;
object-fit: cover;
}
.sent-attachments.images-only {
max-width: min(420px, 100%);
}
.sent-attachments.images-only .sent-image-card {
width: 180px;
height: 180px;
}
.sent-attachments.images-only .sent-image-card img {
object-fit: cover;
background: rgba(var(--v-theme-on-surface), 0.04);
}
.sent-file-card {
width: 220px;
padding: 8px 10px;
}
.sent-attachment-icon {
display: inline-flex;
flex-shrink: 0;
min-width: 34px;
flex-direction: column;
align-items: center;
justify-content: center;
gap: 1px;
}
.sent-attachment-ext {
max-width: 58px;
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
font-size: 10px;
font-weight: 700;
line-height: 12px;
}
.sent-attachment-name {
min-width: 0;
flex: 1;
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
font-size: 13px;
line-height: 18px;
}
.bot-avatar {
margin-top: 2px;
color: rgb(var(--v-theme-primary));
@@ -1335,21 +1097,6 @@ function formatDuration(seconds: number) {
max-width: 82%;
}
.sent-file-card {
width: min(220px, calc(100vw - 28px));
height: 58px;
}
.sent-image-card {
width: 58px;
height: 58px;
}
.sent-attachments.images-only .sent-image-card {
width: min(180px, calc(100vw - 52px));
height: min(180px, calc(100vw - 52px));
}
.message-bubble {
padding: 9px 12px;
}
@@ -1,8 +1,7 @@
<template>
<v-menu v-model="menuOpen" :close-on-content-click="false" location="top" @update:model-value="handleMenuToggle">
<template v-slot:activator="{ props: menuProps }">
<v-chip v-bind="menuProps" class="text-none provider-chip" variant="tonal" :size="chipSize">
<v-chip v-bind="menuProps" class="text-none provider-chip" variant="outlined" size="small">
<v-icon start size="14">mdi-creation</v-icon>
<span v-if="selectedProviderId">
{{ selectedProviderId }}
@@ -61,88 +60,78 @@
</v-card-text>
</v-card>
</v-menu>
</template>
<script setup lang="ts">
import { ref, computed, onMounted } from "vue";
import { useDisplay } from "vuetify";
import axios from "@/utils/request";
import { ref, computed, onMounted } from 'vue';
import axios from 'axios';
interface ModelMetadata {
modalities?: { input?: string[] };
tool_call?: boolean;
reasoning?: boolean;
modalities?: { input?: string[] };
tool_call?: boolean;
reasoning?: boolean;
}
interface ProviderConfig {
id: string;
model: string;
api_base?: string;
model_metadata?: ModelMetadata;
enable?: boolean;
id: string;
model: string;
api_base?: string;
model_metadata?: ModelMetadata;
enable?: boolean;
}
const { mobile } = useDisplay();
const providerConfigs = ref<ProviderConfig[]>([]);
const selectedProviderId = ref("");
const searchQuery = ref("");
const selectedProviderId = ref('');
const searchQuery = ref('');
const menuOpen = ref(false);
const chipSize = computed(() => (mobile.value ? "x-small" : "small"));
const filteredProviders = computed(() => {
if (!searchQuery.value) {
return providerConfigs.value;
}
const query = searchQuery.value.toLowerCase();
return providerConfigs.value.filter(
(p) =>
p.id.toLowerCase().includes(query) ||
p.model.toLowerCase().includes(query),
);
if (!searchQuery.value) {
return providerConfigs.value;
}
const query = searchQuery.value.toLowerCase();
return providerConfigs.value.filter(p =>
p.id.toLowerCase().includes(query) ||
p.model.toLowerCase().includes(query)
);
});
function loadFromStorage() {
const savedProvider = localStorage.getItem("selectedProvider");
if (savedProvider) {
selectedProviderId.value = savedProvider;
}
const savedProvider = localStorage.getItem('selectedProvider');
if (savedProvider) {
selectedProviderId.value = savedProvider;
}
}
function saveToStorage() {
if (selectedProviderId.value) {
localStorage.setItem("selectedProvider", selectedProviderId.value);
}
if (selectedProviderId.value) {
localStorage.setItem('selectedProvider', selectedProviderId.value);
}
}
function loadProviderConfigs() {
axios
.get("/api/config/provider/list", {
params: { provider_type: "chat_completion" },
})
.then((response) => {
if (response.data.status === "ok") {
// 过滤掉 enable 为 false 的配置
providerConfigs.value = (response.data.data || []).filter(
(p: ProviderConfig) => p.enable !== false,
);
}
})
.catch((error) => {
console.error("获取提供商列表失败:", error);
axios.get('/api/config/provider/list', {
params: { provider_type: 'chat_completion' }
}).then(response => {
if (response.data.status === 'ok') {
// 过滤掉 enable 为 false 的配置
providerConfigs.value = (response.data.data || []).filter(
(p: ProviderConfig) => p.enable !== false
);
}
}).catch(error => {
console.error('获取提供商列表失败:', error);
});
}
function selectProvider(provider: ProviderConfig) {
selectedProviderId.value = provider.id;
saveToStorage();
selectedProviderId.value = provider.id;
saveToStorage();
}
function supportsImageInput(provider: ProviderConfig): boolean {
const inputs = provider.model_metadata?.modalities?.input || [];
return inputs.includes("image");
const inputs = provider.model_metadata?.modalities?.input || [];
return inputs.includes('image');
}
function supportsAudioInput(provider: ProviderConfig): boolean {
@@ -151,90 +140,105 @@ function supportsAudioInput(provider: ProviderConfig): boolean {
}
function supportsToolCall(provider: ProviderConfig): boolean {
return Boolean(provider.model_metadata?.tool_call);
return Boolean(provider.model_metadata?.tool_call);
}
function supportsReasoning(provider: ProviderConfig): boolean {
return Boolean(provider.model_metadata?.reasoning);
return Boolean(provider.model_metadata?.reasoning);
}
function getCurrentSelection() {
const provider = providerConfigs.value.find(
(p) => p.id === selectedProviderId.value,
);
return {
providerId: selectedProviderId.value,
modelName: provider?.model || "",
};
const provider = providerConfigs.value.find(p => p.id === selectedProviderId.value);
return {
providerId: selectedProviderId.value,
modelName: provider?.model || ''
};
}
function handleMenuToggle(isOpen: boolean) {
if (isOpen) {
// 每次打开菜单时重新获取数据
loadProviderConfigs();
}
if (isOpen) {
// 每次打开菜单时重新获取数据
loadProviderConfigs();
}
}
onMounted(() => {
loadFromStorage();
loadProviderConfigs();
loadFromStorage();
loadProviderConfigs();
});
defineExpose({
getCurrentSelection,
getCurrentSelection
});
</script>
<style scoped>
.provider-chip {
cursor: pointer;
cursor: pointer;
height: 36px !important;
min-height: 36px !important;
border-color: rgba(var(--v-theme-on-surface), 0.18) !important;
background: transparent !important;
color: rgba(var(--v-theme-on-surface), 0.78) !important;
}
.provider-chip:hover {
border-color: rgba(var(--v-theme-on-surface), 0.34) !important;
background: rgba(var(--v-theme-on-surface), 0.04) !important;
}
.provider-menu-card {
border-radius: 12px !important;
border-radius: 12px !important;
}
.provider-menu-list {
max-height: 280px;
overflow-y: auto;
max-height: 280px;
overflow-y: auto;
}
.provider-menu-item {
margin-bottom: 2px;
border-radius: 8px !important;
min-height: 44px !important;
margin-bottom: 2px;
border-radius: 8px !important;
min-height: 44px !important;
}
.provider-menu-item:hover {
background-color: rgba(103, 58, 183, 0.05);
background-color: rgba(103, 58, 183, 0.05);
}
.provider-menu-item.v-list-item--active {
background-color: rgba(103, 58, 183, 0.1);
background-color: rgba(103, 58, 183, 0.1);
}
.provider-subtitle {
display: flex;
align-items: center;
gap: 8px;
display: flex;
align-items: center;
gap: 8px;
}
.model-name {
font-size: 12px;
color: var(--v-theme-secondaryText);
font-size: 12px;
color: var(--v-theme-secondaryText);
}
.meta-icons {
display: flex;
align-items: center;
gap: 4px;
display: flex;
align-items: center;
gap: 4px;
}
.empty-hint {
font-size: 12px;
color: var(--v-theme-secondaryText);
text-align: center;
padding: 16px;
opacity: 0.6;
font-size: 12px;
color: var(--v-theme-secondaryText);
text-align: center;
padding: 16px;
opacity: 0.6;
}
@media (max-width: 768px) {
.provider-chip {
height: 32px !important;
min-height: 32px !important;
}
}
</style>
+19 -15
View File
@@ -1,12 +1,17 @@
<template>
<v-menu v-bind="$attrs" :close-on-content-click="closeOnContentClick">
<template #activator="{ props: activatorProps }">
<slot name="activator" :props="activatorProps" />
<template v-slot:activator="{ props: activatorProps }">
<slot name="activator" :props="activatorProps"></slot>
</template>
<v-card class="styled-menu-card" elevation="8" rounded="lg">
<v-card
class="styled-menu-card"
:class="{ 'styled-menu-card-borderless': noBorder }"
elevation="8"
rounded="lg"
>
<v-list density="compact" class="styled-menu-list pa-1">
<slot />
<slot></slot>
</v-list>
</v-card>
</v-menu>
@@ -14,17 +19,16 @@
<script setup lang="ts">
defineOptions({
inheritAttrs: false,
});
inheritAttrs: false
})
withDefaults(
defineProps<{
closeOnContentClick?: boolean;
}>(),
{
closeOnContentClick: true,
},
);
withDefaults(defineProps<{
closeOnContentClick?: boolean
noBorder?: boolean
}>(), {
closeOnContentClick: true,
noBorder: false
})
</script>
<style>
File diff suppressed because it is too large Load Diff
@@ -1,151 +1,167 @@
{
"title": "Давай пообщаемся!",
"subtitle": "Общение с AI-помощником",
"input": {
"placeholder": "Введите сообщение...",
"send": "Отправить",
"clear": "Очистить",
"upload": "Загрузить файл",
"voice": "Голосовой ввод",
"recordingPrompt": "Запись... говорите",
"chatPrompt": "Давай пообщаемся!",
"dropToUpload": "Отпустите, чтобы загрузить файл",
"stopGenerating": "Остановить генерацию"
},
"message": {
"user": "Вы",
"assistant": "Ассистент",
"system": "Система",
"error": "Ошибка в сообщении",
"loading": "Думаю..."
},
"voice": {
"start": "Начать запись",
"stop": "Стоп",
"recording": "Запись",
"processing": "Обработка...",
"error": "Ошибка записи",
"listening": "Слушаю...",
"speaking": "Говорю",
"startRecording": "Начать голосовой ввод",
"liveMode": "Общение в реальном времени"
},
"welcome": {
"title": "Добро пожаловать в AstrBot",
"subtitle": "Ваш умный помощник",
"quickActions": "Быстрые действия",
"examples": "Примеры вопросов"
},
"actions": {
"copy": "Копировать",
"regenerate": "Перегенерировать",
"like": "Нравится",
"dislike": "Не нравится",
"share": "Поделиться",
"newChat": "Новый чат",
"deleteChat": "Удалить чат",
"editTitle": "Изменить заголовок",
"fullscreen": "На весь экран",
"exitFullscreen": "Выход из полноэкранного режима",
"reply": "Ответить",
"providerConfig": "Настройки AI",
"toolsUsed": "Использованные инструменты",
"toolCallUsed": "Использован инструмент {name}",
"pythonCodeAnalysis": "Использован анализ кода Python"
},
"ipython": {
"output": "Вывод"
},
"conversation": {
"newConversation": "Новый чат",
"noHistory": "История диалогов пуста",
"systemStatus": "Статус системы",
"llmService": "Сервис LLM",
"speechToText": "Преобразование речи",
"editDisplayName": "Изменить имя чата",
"displayName": "Имя чата",
"displayNameUpdated": "Имя чата обновлено",
"displayNameUpdateFailed": "Не удалось обновить имя чата",
"confirmDelete": "Вы уверены, что хотите удалить «{name}»? Это действие необратимо."
},
"modes": {
"darkMode": "Темная тема",
"lightMode": "Светлая тема"
},
"shortcuts": {
"help": "Справка",
"voiceRecord": "Запись голоса",
"pasteImage": "Вставить изображение",
"sendKey": {
"title": "Клавиша отправки",
"enterToSend": "Enter для отправки",
"shiftEnterToSend": "Shift+Enter для отправки"
}
},
"streaming": {
"enabled": "Потоковый ответ включен",
"disabled": "Потоковый ответ выключен",
"on": "Поток",
"off": "Обычный"
},
"transport": {
"title": "Протокол передачи",
"sse": "SSE",
"websocket": "WebSocket"
},
"config": {
"title": "Конфигурация"
},
"reasoning": {
"thinking": "Рассуждение"
},
"reply": {
"replyTo": "В ответ на",
"notFound": "Сообщение не найдено"
},
"project": {
"title": "Проект",
"create": "Создать проект",
"edit": "Изменить проект",
"name": "Имя проекта",
"emoji": "Иконка (Emoji)",
"description": "Описание проекта (опционально)",
"noSessions": "В этом проекте пока нет диалогов",
"confirmDelete": "Вы уверены, что хотите удалить проект «{title}»? Диалоги внутри проекта не будут удалены."
},
"time": {
"today": "Сегодня",
"yesterday": "Вчера"
},
"stats": {
"tokens": "Токены",
"inputTokens": "Входящие",
"outputTokens": "Исходящие",
"cachedTokens": "Кэшированные",
"duration": "Время",
"ttft": "Время до первого токена"
},
"refs": {
"title": "Ссылки",
"sources": "Источники"
},
"connection": {
"title": "Статус подключения",
"message": "Системе необходимо переустановить соединение с чатом.",
"reasons": "Это может быть вызвано следующими причинами:",
"reasonWindowResize": "Изменение размера окна (нормально)",
"reasonMultipleTabs": "Страница чата открыта в другой вкладке",
"reasonNetworkIssue": "Временная проблема с сетью",
"notice": "Примечание: для стабильной работы допускается только одно активное соединение. Если вы используете чат в нескольких вкладках, рекомендуем оставить только одну.",
"understand": "Понятно",
"status": {
"reconnecting": "Переподключение...",
"reconnected": "Соединение восстановлено",
"failed": "Ошибка подключения, обновите страницу"
}
},
"errors": {
"sendMessageFailed": "Ошибка отправки сообщения, попробуйте еще раз",
"createSessionFailed": "Ошибка создания сессии, обновите страницу"
{
"title": "Давай пообщаемся!",
"subtitle": "Общение с AI-помощником",
"input": {
"placeholder": "Введите сообщение...",
"send": "Отправить",
"clear": "Очистить",
"upload": "Загрузить файл",
"voice": "Голосовой ввод",
"recordingPrompt": "Запись... говорите",
"chatPrompt": "Давай пообщаемся!",
"dropToUpload": "Отпустите, чтобы загрузить файл",
"stopGenerating": "Остановить генерацию"
},
"message": {
"user": "Вы",
"assistant": "Ассистент",
"system": "Система",
"error": "Ошибка в сообщении",
"loading": "Думаю..."
},
"voice": {
"start": "Начать запись",
"stop": "Стоп",
"recording": "Запись",
"processing": "Обработка...",
"error": "Ошибка записи",
"listening": "Слушаю...",
"speaking": "Говорю",
"startRecording": "Начать голосовой ввод",
"liveMode": "Общение в реальном времени"
},
"welcome": {
"title": "Добро пожаловать в AstrBot",
"subtitle": "Ваш умный помощник",
"quickActions": "Быстрые действия",
"examples": "Примеры вопросов"
},
"actions": {
"copy": "Копировать",
"regenerate": "Перегенерировать",
"retry": "Повторить",
"retryWithModel": "Использовать другую модель",
"noAvailableModels": "Нет доступных моделей",
"like": "Нравится",
"dislike": "Не нравится",
"share": "Поделиться",
"newChat": "Новый чат",
"deleteChat": "Удалить чат",
"editTitle": "Изменить заголовок",
"fullscreen": "На весь экран",
"exitFullscreen": "Выход из полноэкранного режима",
"reply": "Ответить",
"providerConfig": "Настройки AI",
"toolsUsed": "Использованные инструменты",
"toolCallUsed": "Использован инструмент {name}",
"pythonCodeAnalysis": "Использован анализ кода Python"
},
"ipython": {
"output": "Вывод"
},
"toolStatus": {
"done": "Готово",
"running": "Выполняется"
},
"thread": {
"title": "Thread",
"askInThread": "Спросить в ветке",
"count": "Веток: {count}",
"createFailed": "Не удалось создать ветку",
"delete": "Удалить ветку",
"confirmDelete": "Удалить эту ветку? Это действие необратимо.",
"placeholder": "Спросить об этом фрагменте..."
},
"conversation": {
"newConversation": "Новый чат",
"noHistory": "История диалогов пуста",
"systemStatus": "Статус системы",
"llmService": "Сервис LLM",
"speechToText": "Преобразование речи",
"editDisplayName": "Изменить имя чата",
"displayName": "Имя чата",
"displayNameUpdated": "Имя чата обновлено",
"displayNameUpdateFailed": "Не удалось обновить имя чата",
"confirmDelete": "Вы уверены, что хотите удалить «{name}»? Это действие необратимо."
},
"modes": {
"darkMode": "Темная тема",
"lightMode": "Светлая тема"
},
"shortcuts": {
"help": "Справка",
"voiceRecord": "Запись голоса",
"pasteImage": "Вставить изображение",
"sendKey": {
"title": "Клавиша отправки",
"enterToSend": "Enter для отправки",
"shiftEnterToSend": "Shift+Enter для отправки"
}
},
"streaming": {
"enabled": "Потоковый ответ включен",
"disabled": "Потоковый ответ выключен",
"on": "Поток",
"off": "Обычный"
},
"transport": {
"title": "Протокол передачи",
"sse": "SSE",
"websocket": "WebSocket"
},
"config": {
"title": "Конфигурация"
},
"reasoning": {
"thinking": "Рассуждение"
},
"reply": {
"replyTo": "В ответ на",
"notFound": "Сообщение не найдено"
},
"project": {
"title": "Проект",
"create": "Создать проект",
"edit": "Изменить проект",
"name": "Имя проекта",
"emoji": "Иконка (Emoji)",
"description": "Описание проекта (опционально)",
"noSessions": "В этом проекте пока нет диалогов",
"confirmDelete": "Вы уверены, что хотите удалить проект «{title}»? Диалоги внутри проекта не будут удалены."
},
"time": {
"today": "Сегодня",
"yesterday": "Вчера"
},
"stats": {
"tokens": "Токены",
"inputTokens": "Входящие",
"outputTokens": "Исходящие",
"cachedTokens": "Кэшированные",
"duration": "Время",
"ttft": "Время до первого токена"
},
"refs": {
"title": "Ссылки",
"sources": "Источники"
},
"connection": {
"title": "Статус подключения",
"message": "Системе необходимо переустановить соединение с чатом.",
"reasons": "Это может быть вызвано следующими причинами:",
"reasonWindowResize": "Изменение размера окна (нормально)",
"reasonMultipleTabs": "Страница чата открыта в другой вкладке",
"reasonNetworkIssue": "Временная проблема с сетью",
"notice": "Примечание: для стабильной работы допускается только одно активное соединение. Если вы используете чат в нескольких вкладках, рекомендуем оставить только одну.",
"understand": "Понятно",
"status": {
"reconnecting": "Переподключение...",
"reconnected": "Соединение восстановлено",
"failed": "Ошибка подключения, обновите страницу"
}
},
"errors": {
"sendMessageFailed": "Ошибка отправки сообщения, попробуйте еще раз",
"createSessionFailed": "Ошибка создания сессии, обновите страницу"
}
}