From 716e651d80298106676158a5d73ed7afe7ff58b2 Mon Sep 17 00:00:00 2001 From: LIghtJUNction Date: Wed, 29 Apr 2026 03:10:03 +0800 Subject: [PATCH] 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 --- astrbot/core/agent/message.py | 64 +- .../agent/runners/tool_loop_agent_runner.py | 281 +- astrbot/core/astr_main_agent.py | 139 +- astrbot/core/db/sqlite.py | 1203 +++++---- .../method/agent_sub_stages/internal.py | 29 +- .../sources/webchat/webchat_adapter.py | 47 +- astrbot/dashboard/routes/chat.py | 773 +++--- astrbot/dashboard/routes/live_chat.py | 406 ++- dashboard/src/components/chat/Chat.vue | 2273 +++++++++++------ dashboard/src/components/chat/ChatInput.vue | 290 +-- .../src/components/chat/ChatMessageList.vue | 255 +- .../src/components/chat/ProviderModelMenu.vue | 188 +- .../src/components/shared/StyledMenu.vue | 34 +- dashboard/src/composables/useMessages.ts | 1861 ++++++-------- .../src/i18n/locales/ru-RU/features/chat.json | 314 +-- 15 files changed, 4005 insertions(+), 4152 deletions(-) diff --git a/astrbot/core/agent/message.py b/astrbot/core/agent/message.py index 1d60845ba..ad3b57cb2 100644 --- a/astrbot/core/agent/message.py +++ b/astrbot/core/agent/message.py @@ -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 diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index ebecd3d03..132ca8192 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -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: diff --git a/astrbot/core/astr_main_agent.py b/astrbot/core/astr_main_agent.py index 54f50b628..a9ddb2e7b 100644 --- a/astrbot/core/astr_main_agent.py +++ b/astrbot/core/astr_main_agent.py @@ -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"{caption}"), + TextPart(text=f"{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"\n{quoted_content}\n" @@ -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 "" 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"{thread_selected_text.strip()}" - ), - ), + ) + ) ) 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() diff --git a/astrbot/core/db/sqlite.py b/astrbot/core/db/sqlite.py index 2de471671..d79ac9d70 100644 --- a/astrbot/core/db/sqlite.py +++ b/astrbot/core/db/sqlite.py @@ -1,10 +1,10 @@ import asyncio import threading -from collections.abc import Awaitable, Callable, Sequence +import typing as T +from collections.abc import Awaitable, Callable from datetime import datetime, timedelta, timezone -from typing import Any, TypeVar -from sqlalchemy import Row +from sqlalchemy import CursorResult, Row from sqlalchemy.ext.asyncio import AsyncSession from sqlmodel import col, delete, desc, func, or_, select, text, update @@ -28,11 +28,15 @@ from astrbot.core.db.po import ( SQLModel, WebChatThread, ) -from astrbot.core.db.po import Platform as DeprecatedPlatformStat -from astrbot.core.db.po import Stats as DeprecatedStats +from astrbot.core.db.po import ( + Platform as DeprecatedPlatformStat, +) +from astrbot.core.db.po import ( + Stats as DeprecatedStats, +) from astrbot.core.sentinels import NOT_GIVEN -TxResult = TypeVar("TxResult") +TxResult = T.TypeVar("TxResult") CRON_FIELD_NOT_SET = object() @@ -53,6 +57,7 @@ class SQLiteDatabase(BaseDatabase): await conn.execute(text("PRAGMA temp_store=MEMORY")) await conn.execute(text("PRAGMA mmap_size=134217728")) await conn.execute(text("PRAGMA optimize")) + # 确保 personas 表有 folder_id、sort_order、skills 列(前向兼容) await self._ensure_persona_folder_columns(conn) await self._ensure_persona_skills_column(conn) await self._ensure_persona_custom_error_message_column(conn) @@ -60,42 +65,45 @@ class SQLiteDatabase(BaseDatabase): await conn.commit() async def _ensure_persona_folder_columns(self, conn) -> None: - """确保 personas 表有 folder_id 和 sort_order 列。 + """确保 personas 表有 folder_id 和 sort_order 列。 - 这是为了支持旧版数据库的平滑升级。新版数据库通过 SQLModel - 的 metadata.create_all 自动创建这些列。 + 这是为了支持旧版数据库的平滑升级。新版数据库通过 SQLModel + 的 metadata.create_all 自动创建这些列。 """ result = await conn.execute(text("PRAGMA table_info(personas)")) columns = {row[1] for row in result.fetchall()} + if "folder_id" not in columns: await conn.execute( text( - "ALTER TABLE personas ADD COLUMN folder_id VARCHAR(36) DEFAULT NULL", - ), + "ALTER TABLE personas ADD COLUMN folder_id VARCHAR(36) DEFAULT NULL" + ) ) if "sort_order" not in columns: await conn.execute( - text("ALTER TABLE personas ADD COLUMN sort_order INTEGER DEFAULT 0"), + text("ALTER TABLE personas ADD COLUMN sort_order INTEGER DEFAULT 0") ) async def _ensure_persona_skills_column(self, conn) -> None: - """确保 personas 表有 skills 列。 + """确保 personas 表有 skills 列。 - 这是为了支持旧版数据库的平滑升级。新版数据库通过 SQLModel - 的 metadata.create_all 自动创建这些列。 + 这是为了支持旧版数据库的平滑升级。新版数据库通过 SQLModel + 的 metadata.create_all 自动创建这些列。 """ result = await conn.execute(text("PRAGMA table_info(personas)")) columns = {row[1] for row in result.fetchall()} + if "skills" not in columns: await conn.execute(text("ALTER TABLE personas ADD COLUMN skills JSON")) async def _ensure_persona_custom_error_message_column(self, conn) -> None: - """确保 personas 表有 custom_error_message 列。""" + """确保 personas 表有 custom_error_message 列。""" result = await conn.execute(text("PRAGMA table_info(personas)")) columns = {row[1] for row in result.fetchall()} + if "custom_error_message" not in columns: await conn.execute( - text("ALTER TABLE personas ADD COLUMN custom_error_message TEXT"), + text("ALTER TABLE personas ADD COLUMN custom_error_message TEXT") ) async def _ensure_platform_message_history_checkpoint_column(self, conn) -> None: @@ -107,15 +115,15 @@ class SQLiteDatabase(BaseDatabase): await conn.execute( text( "ALTER TABLE platform_message_history " - "ADD COLUMN llm_checkpoint_id VARCHAR DEFAULT NULL", - ), + "ADD COLUMN llm_checkpoint_id VARCHAR DEFAULT NULL" + ) ) await conn.execute( text( "CREATE INDEX IF NOT EXISTS " "ix_platform_message_history_llm_checkpoint_id " - "ON platform_message_history (llm_checkpoint_id)", - ), + "ON platform_message_history (llm_checkpoint_id)" + ) ) # ==== @@ -131,6 +139,7 @@ class SQLiteDatabase(BaseDatabase): ) -> None: """Insert a new platform statistic record.""" async with self.get_db() as session: + session: AsyncSession async with session.begin(): if timestamp is None: timestamp = datetime.now().replace( @@ -140,9 +149,12 @@ class SQLiteDatabase(BaseDatabase): ) current_hour = timestamp await session.execute( - text( - "\n INSERT INTO platform_stats (timestamp, platform_id, platform_type, count)\n VALUES (:timestamp, :platform_id, :platform_type, :count)\n ON CONFLICT(timestamp, platform_id, platform_type) DO UPDATE SET\n count = platform_stats.count + EXCLUDED.count\n ", - ), + text(""" + INSERT INTO platform_stats (timestamp, platform_id, platform_type, count) + VALUES (:timestamp, :platform_id, :platform_type, :count) + ON CONFLICT(timestamp, platform_id, platform_type) DO UPDATE SET + count = platform_stats.count + EXCLUDED.count + """), { "timestamp": current_hour, "platform_id": platform_id, @@ -154,6 +166,7 @@ class SQLiteDatabase(BaseDatabase): async def count_platform_stats(self) -> int: """Count the number of platform statistics records.""" async with self.get_db() as session: + session: AsyncSession result = await session.execute( select(func.count(col(PlatformStat.platform_id))).select_from( PlatformStat, @@ -165,12 +178,16 @@ class SQLiteDatabase(BaseDatabase): async def get_platform_stats(self, offset_sec: int = 86400) -> list[PlatformStat]: """Get platform statistics within the specified offset in seconds and group by platform_id.""" async with self.get_db() as session: + session: AsyncSession now = datetime.now() start_time = now - timedelta(seconds=offset_sec) result = await session.execute( - text( - "\n SELECT * FROM platform_stats\n WHERE timestamp >= :start_time\n GROUP BY platform_id\n ORDER BY timestamp DESC\n ", - ), + text(""" + SELECT * FROM platform_stats + WHERE timestamp >= :start_time + GROUP BY platform_id + ORDER BY timestamp DESC + """), {"start_time": start_time}, ) return list(result.scalars().all()) @@ -189,51 +206,66 @@ class SQLiteDatabase(BaseDatabase): """Insert a provider stat record for a single agent response.""" stats = stats or {} token_usage = stats.get("token_usage", {}) + token_input_other = int(token_usage.get("input_other", 0) or 0) token_input_cached = int(token_usage.get("input_cached", 0) or 0) token_output = int(token_usage.get("output", 0) or 0) + start_time = float(stats.get("start_time", 0.0) or 0.0) end_time = float(stats.get("end_time", 0.0) or 0.0) time_to_first_token = float(stats.get("time_to_first_token", 0.0) or 0.0) - async with self.get_db() as session, session.begin(): - record = ProviderStat( - agent_type=agent_type, - status=status, - umo=umo, - conversation_id=conversation_id, - provider_id=provider_id, - provider_model=provider_model, - token_input_other=token_input_other, - token_input_cached=token_input_cached, - token_output=token_output, - start_time=start_time, - end_time=end_time, - time_to_first_token=time_to_first_token, - ) - session.add(record) - await session.flush() - await session.refresh(record) - return record + + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + record = ProviderStat( + agent_type=agent_type, + status=status, + umo=umo, + conversation_id=conversation_id, + provider_id=provider_id, + provider_model=provider_model, + token_input_other=token_input_other, + token_input_cached=token_input_cached, + token_output=token_output, + start_time=start_time, + end_time=end_time, + time_to_first_token=time_to_first_token, + ) + session.add(record) + await session.flush() + await session.refresh(record) + return record + + # ==== + # Conversation Management + # ==== async def get_conversations(self, user_id=None, platform_id=None): async with self.get_db() as session: + session: AsyncSession query = select(ConversationV2) + if user_id: query = query.where(ConversationV2.user_id == user_id) if platform_id: query = query.where(ConversationV2.platform_id == platform_id) + # order by query = query.order_by(desc(ConversationV2.created_at)) result = await session.execute(query) + return result.scalars().all() async def get_conversation_by_id(self, cid): async with self.get_db() as session: + session: AsyncSession query = select(ConversationV2).where(ConversationV2.conversation_id == cid) result = await session.execute(query) return result.scalar_one_or_none() async def get_all_conversations(self, page=1, page_size=20): async with self.get_db() as session: + session: AsyncSession offset = (page - 1) * page_size result = await session.execute( select(ConversationV2) @@ -252,7 +284,10 @@ class SQLiteDatabase(BaseDatabase): **kwargs, ): async with self.get_db() as session: + session: AsyncSession + # Build the base query with filters base_query = select(ConversationV2) + if platform_ids: base_query = base_query.where( col(ConversationV2.platform_id).in_(platform_ids), @@ -276,9 +311,13 @@ class SQLiteDatabase(BaseDatabase): base_query = base_query.where( col(ConversationV2.platform_id).in_(kwargs["platforms"]), ) + + # Get total count matching the filters count_query = select(func.count()).select_from(base_query.subquery()) total_count = await session.execute(count_query) total = total_count.scalar_one() + + # Get paginated results offset = (page - 1) * page_size result_query = ( base_query.order_by(desc(ConversationV2.created_at)) @@ -287,7 +326,8 @@ class SQLiteDatabase(BaseDatabase): ) result = await session.execute(result_query) conversations = result.scalars().all() - return (conversations, total) + + return conversations, total async def create_conversation( self, @@ -307,61 +347,63 @@ class SQLiteDatabase(BaseDatabase): kwargs["created_at"] = created_at if updated_at: kwargs["updated_at"] = updated_at - async with self.get_db() as session, session.begin(): - new_conversation = ConversationV2( - user_id=user_id, - content=content or [], - platform_id=platform_id, - title=title, - persona_id=persona_id, - **kwargs, - ) - session.add(new_conversation) - return new_conversation + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + new_conversation = ConversationV2( + user_id=user_id, + content=content or [], + platform_id=platform_id, + title=title, + persona_id=persona_id, + **kwargs, + ) + session.add(new_conversation) + return new_conversation async def update_conversation( - self, - cid, - title=None, - persona_id=None, - clear_persona: bool = False, - content=None, - token_usage=None, + self, cid, title=None, persona_id=None, content=None, token_usage=None ): - async with self.get_db() as session, session.begin(): - query = update(ConversationV2).where( - col(ConversationV2.conversation_id) == cid, - ) - values = {} - if title is not None: - values["title"] = title - if clear_persona: - values["persona_id"] = None - elif persona_id is not None: - values["persona_id"] = persona_id - if content is not None: - values["content"] = content - if token_usage is not None: - values["token_usage"] = token_usage - if not values: - return None - query = query.values(**values) - await session.execute(query) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + query = update(ConversationV2).where( + col(ConversationV2.conversation_id) == cid, + ) + values = {} + if title is not None: + values["title"] = title + if persona_id is not None: + values["persona_id"] = persona_id + if content is not None: + values["content"] = content + if token_usage is not None: + values["token_usage"] = token_usage + if not values: + return None + query = query.values(**values) + await session.execute(query) return await self.get_conversation_by_id(cid) async def delete_conversation(self, cid) -> None: - async with self.get_db() as session, session.begin(): - await session.execute( - delete(ConversationV2).where( - col(ConversationV2.conversation_id) == cid, - ), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + await session.execute( + delete(ConversationV2).where( + col(ConversationV2.conversation_id) == cid, + ), + ) async def delete_conversations_by_user_id(self, user_id: str) -> None: - async with self.get_db() as session, session.begin(): - await session.execute( - delete(ConversationV2).where(col(ConversationV2.user_id) == user_id), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + await session.execute( + delete(ConversationV2).where( + col(ConversationV2.user_id) == user_id + ), + ) async def get_session_conversations( self, @@ -372,13 +414,15 @@ class SQLiteDatabase(BaseDatabase): ) -> tuple[list[dict], int]: """Get paginated session conversations with joined conversation and persona details.""" async with self.get_db() as session: + session: AsyncSession offset = (page - 1) * page_size + base_query = ( select( col(Preference.scope_id).label("session_id"), func.json_extract(Preference.value, "$.val").label( "conversation_id", - ), + ), # type: ignore col(ConversationV2.persona_id).label("persona_id"), col(ConversationV2.title).label("title"), col(Persona.persona_id).label("persona_name"), @@ -395,6 +439,8 @@ class SQLiteDatabase(BaseDatabase): ) .where(Preference.scope == "umo", Preference.key == "sel_conv_id") ) + + # 搜索筛选 if search_query: search_pattern = f"%{search_query}%" base_query = base_query.where( @@ -404,15 +450,23 @@ class SQLiteDatabase(BaseDatabase): col(Persona.persona_id).ilike(search_pattern), ), ) + + # 平台筛选 if platform: platform_pattern = f"{platform}:%" base_query = base_query.where( col(Preference.scope_id).like(platform_pattern), ) + + # 排序 base_query = base_query.order_by(Preference.scope_id) + + # 分页结果 result_query = base_query.offset(offset).limit(page_size) result = await session.execute(result_query) rows = result.fetchall() + + # 查询总数(应用相同的筛选条件) count_base_query = ( select(func.count(col(Preference.scope_id))) .select_from(Preference) @@ -427,6 +481,8 @@ class SQLiteDatabase(BaseDatabase): ) .where(Preference.scope == "umo", Preference.key == "sel_conv_id") ) + + # 应用相同的搜索和平台筛选条件到计数查询 if search_query: search_pattern = f"%{search_query}%" count_base_query = count_base_query.where( @@ -436,13 +492,16 @@ class SQLiteDatabase(BaseDatabase): col(Persona.persona_id).ilike(search_pattern), ), ) + if platform: platform_pattern = f"{platform}:%" count_base_query = count_base_query.where( col(Preference.scope_id).like(platform_pattern), ) + total_result = await session.execute(count_base_query) total = total_result.scalar() or 0 + sessions_data = [ { "session_id": row.session_id, @@ -453,7 +512,7 @@ class SQLiteDatabase(BaseDatabase): } for row in rows ] - return (sessions_data, total) + return sessions_data, total async def insert_platform_message_history( self, @@ -500,7 +559,7 @@ class SQLiteDatabase(BaseDatabase): await session.execute( update(PlatformMessageHistory) .where(PlatformMessageHistory.id == message_id) - .values(**values), + .values(**values) ) async def delete_platform_message_history_by_id(self, message_id: int) -> None: @@ -510,8 +569,8 @@ class SQLiteDatabase(BaseDatabase): async with session.begin(): await session.execute( delete(PlatformMessageHistory).where( - PlatformMessageHistory.id == message_id, - ), + PlatformMessageHistory.id == message_id + ) ) async def delete_platform_message_offset( @@ -521,16 +580,18 @@ class SQLiteDatabase(BaseDatabase): offset_sec=86400, ) -> None: """Delete platform message history records newer than the specified offset.""" - async with self.get_db() as session, session.begin(): - now = datetime.now() - cutoff_time = now - timedelta(seconds=offset_sec) - await session.execute( - delete(PlatformMessageHistory).where( - col(PlatformMessageHistory.platform_id) == platform_id, - col(PlatformMessageHistory.user_id) == user_id, - col(PlatformMessageHistory.created_at) >= cutoff_time, - ), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + now = datetime.now() + cutoff_time = now - timedelta(seconds=offset_sec) + await session.execute( + delete(PlatformMessageHistory).where( + col(PlatformMessageHistory.platform_id) == platform_id, + col(PlatformMessageHistory.user_id) == user_id, + col(PlatformMessageHistory.created_at) >= cutoff_time, + ), + ) async def get_platform_message_history( self, @@ -541,6 +602,7 @@ class SQLiteDatabase(BaseDatabase): ): """Get platform message history records.""" async with self.get_db() as session: + session: AsyncSession offset = (page - 1) * page_size query = ( select(PlatformMessageHistory) @@ -553,108 +615,14 @@ class SQLiteDatabase(BaseDatabase): result = await session.execute(query.offset(offset).limit(page_size)) return result.scalars().all() - async def list_sdk_platform_message_history( - self, - platform_id, - user_id, - cursor_id=None, - limit=50, - include_total=False, - ): - """List SDK message history records ordered by descending id.""" - async with self.get_db() as session: - query = ( - select(PlatformMessageHistory) - .where( - PlatformMessageHistory.platform_id == platform_id, - PlatformMessageHistory.user_id == user_id, - ) - .order_by(desc(PlatformMessageHistory.id)) - ) - if cursor_id is not None: - query = query.where(PlatformMessageHistory.id < cursor_id) - result = await session.execute(query.limit(limit)) - total: int | None = None - if include_total: - total_query = ( - select(func.count()) - .select_from(PlatformMessageHistory) - .where( - PlatformMessageHistory.platform_id == platform_id, - PlatformMessageHistory.user_id == user_id, - ) - ) - total_result = await session.execute(total_query) - total = int(total_result.scalar() or 0) - return (list(result.scalars().all()), total) - - async def delete_platform_message_before(self, platform_id, user_id, before) -> int: - """Delete platform message history records strictly older than the boundary.""" - async with self.get_db() as session, session.begin(): - result = await session.execute( - delete(PlatformMessageHistory).where( - col(PlatformMessageHistory.platform_id) == platform_id, - col(PlatformMessageHistory.user_id) == user_id, - col(PlatformMessageHistory.created_at) < before, - ), - ) - return int(getattr(result, "rowcount", 0) or 0) - - async def delete_platform_message_after(self, platform_id, user_id, after) -> int: - """Delete platform message history records strictly newer than the boundary.""" - async with self.get_db() as session, session.begin(): - result = await session.execute( - delete(PlatformMessageHistory).where( - col(PlatformMessageHistory.platform_id) == platform_id, - col(PlatformMessageHistory.user_id) == user_id, - col(PlatformMessageHistory.created_at) > after, - ), - ) - return int(getattr(result, "rowcount", 0) or 0) - - async def delete_all_platform_message_history(self, platform_id, user_id) -> int: - """Delete all platform message history records for a specific user.""" - async with self.get_db() as session, session.begin(): - result = await session.execute( - delete(PlatformMessageHistory).where( - col(PlatformMessageHistory.platform_id) == platform_id, - col(PlatformMessageHistory.user_id) == user_id, - ), - ) - return int(getattr(result, "rowcount", 0) or 0) - - async def find_platform_message_history_by_idempotency_key( - self, - platform_id, - user_id, - idempotency_key, - ) -> PlatformMessageHistory | None: - """Find a SDK message history record by its idempotency key.""" - async with self.get_db() as session: - query = ( - select(PlatformMessageHistory) - .where( - PlatformMessageHistory.platform_id == platform_id, - PlatformMessageHistory.user_id == user_id, - func.json_extract( - PlatformMessageHistory.content, - "$.idempotency_key", - ) - == str(idempotency_key), - ) - .order_by(desc(PlatformMessageHistory.id)) - ) - result = await session.execute(query.limit(1)) - return result.scalar_one_or_none() - async def get_platform_message_history_by_id( - self, - message_id: int, + self, message_id: int ) -> PlatformMessageHistory | None: """Get a platform message history record by its ID.""" async with self.get_db() as session: + session: AsyncSession query = select(PlatformMessageHistory).where( - PlatformMessageHistory.id == message_id, + PlatformMessageHistory.id == message_id ) result = await session.execute(query) return result.scalar_one_or_none() @@ -691,7 +659,7 @@ class SQLiteDatabase(BaseDatabase): async with self.get_db() as session: session: AsyncSession result = await session.execute( - select(WebChatThread).where(WebChatThread.thread_id == thread_id), + select(WebChatThread).where(WebChatThread.thread_id == thread_id) ) return result.scalar_one_or_none() @@ -704,7 +672,7 @@ class SQLiteDatabase(BaseDatabase): async with self.get_db() as session: session: AsyncSession query = select(WebChatThread).where( - WebChatThread.parent_session_id == parent_session_id, + WebChatThread.parent_session_id == parent_session_id ) if creator is not None: query = query.where(WebChatThread.creator == creator) @@ -738,7 +706,7 @@ class SQLiteDatabase(BaseDatabase): session: AsyncSession async with session.begin(): await session.execute( - delete(WebChatThread).where(WebChatThread.thread_id == thread_id), + delete(WebChatThread).where(WebChatThread.thread_id == thread_id) ) async def delete_webchat_threads_by_parent_session( @@ -755,8 +723,8 @@ class SQLiteDatabase(BaseDatabase): async with session.begin(): await session.execute( delete(WebChatThread).where( - col(WebChatThread.thread_id).in_(thread_ids), - ), + col(WebChatThread.thread_id).in_(thread_ids) + ) ) return thread_ids @@ -774,7 +742,7 @@ class SQLiteDatabase(BaseDatabase): select(WebChatThread.thread_id).where( WebChatThread.parent_session_id == parent_session_id, col(WebChatThread.parent_message_id).in_(parent_message_ids), - ), + ) ) thread_ids = list(result.scalars().all()) if not thread_ids: @@ -784,21 +752,28 @@ class SQLiteDatabase(BaseDatabase): async with session.begin(): await session.execute( delete(WebChatThread).where( - col(WebChatThread.thread_id).in_(thread_ids), - ), + col(WebChatThread.thread_id).in_(thread_ids) + ) ) return thread_ids async def insert_attachment(self, path, type, mime_type): """Insert a new attachment record.""" - async with self.get_db() as session, session.begin(): - new_attachment = Attachment(path=path, type=type, mime_type=mime_type) - session.add(new_attachment) - return new_attachment + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + new_attachment = Attachment( + path=path, + type=type, + mime_type=mime_type, + ) + session.add(new_attachment) + return new_attachment async def get_attachment_by_id(self, attachment_id): """Get an attachment by its ID.""" async with self.get_db() as session: + session: AsyncSession query = select(Attachment).where(Attachment.attachment_id == attachment_id) result = await session.execute(query) return result.scalar_one_or_none() @@ -808,8 +783,9 @@ class SQLiteDatabase(BaseDatabase): if not attachment_ids: return [] async with self.get_db() as session: + session: AsyncSession query = select(Attachment).where( - col(Attachment.attachment_id).in_(attachment_ids), + col(Attachment.attachment_id).in_(attachment_ids) ) result = await session.execute(query) return list(result.scalars().all()) @@ -819,12 +795,14 @@ class SQLiteDatabase(BaseDatabase): Returns True if the attachment was deleted, False if it was not found. """ - async with self.get_db() as session, session.begin(): - query = delete(Attachment).where( - col(Attachment.attachment_id) == attachment_id, - ) - result = await session.execute(query) - return getattr(result, "rowcount", 0) > 0 + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + query = delete(Attachment).where( + col(Attachment.attachment_id) == attachment_id + ) + result = T.cast(CursorResult, await session.execute(query)) + return result.rowcount > 0 async def delete_attachments(self, attachment_ids: list[str]) -> int: """Delete multiple attachments by their IDs. @@ -833,12 +811,14 @@ class SQLiteDatabase(BaseDatabase): """ if not attachment_ids: return 0 - async with self.get_db() as session, session.begin(): - query = delete(Attachment).where( - col(Attachment.attachment_id).in_(attachment_ids), - ) - result = await session.execute(query) - return getattr(result, "rowcount", 0) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + query = delete(Attachment).where( + col(Attachment.attachment_id).in_(attachment_ids) + ) + result = T.cast(CursorResult, await session.execute(query)) + return result.rowcount async def create_api_key( self, @@ -850,39 +830,44 @@ class SQLiteDatabase(BaseDatabase): expires_at: datetime | None = None, ) -> ApiKey: """Create a new API key record.""" - async with self.get_db() as session, session.begin(): - api_key = ApiKey( - name=name, - key_hash=key_hash, - key_prefix=key_prefix, - scopes=scopes, - created_by=created_by, - expires_at=expires_at, - ) - session.add(api_key) - await session.flush() - await session.refresh(api_key) - return api_key + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + api_key = ApiKey( + name=name, + key_hash=key_hash, + key_prefix=key_prefix, + scopes=scopes, + created_by=created_by, + expires_at=expires_at, + ) + session.add(api_key) + await session.flush() + await session.refresh(api_key) + return api_key async def list_api_keys(self) -> list[ApiKey]: """List all API keys.""" async with self.get_db() as session: + session: AsyncSession result = await session.execute( - select(ApiKey).order_by(desc(ApiKey.created_at)), + select(ApiKey).order_by(desc(ApiKey.created_at)) ) return list(result.scalars().all()) async def get_api_key_by_id(self, key_id: str) -> ApiKey | None: """Get an API key by key_id.""" async with self.get_db() as session: + session: AsyncSession result = await session.execute( - select(ApiKey).where(ApiKey.key_id == key_id), + select(ApiKey).where(ApiKey.key_id == key_id) ) return result.scalar_one_or_none() async def get_active_api_key_by_hash(self, key_hash: str) -> ApiKey | None: """Get an active API key by hash (not revoked, not expired).""" async with self.get_db() as session: + session: AsyncSession now = datetime.now(timezone.utc) query = select(ApiKey).where( ApiKey.key_hash == key_hash, @@ -894,31 +879,40 @@ class SQLiteDatabase(BaseDatabase): async def touch_api_key(self, key_id: str) -> None: """Update last_used_at of an API key.""" - async with self.get_db() as session, session.begin(): - await session.execute( - update(ApiKey) - .where(col(ApiKey.key_id) == key_id) - .values(last_used_at=datetime.now(timezone.utc)), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + await session.execute( + update(ApiKey) + .where(col(ApiKey.key_id) == key_id) + .values(last_used_at=datetime.now(timezone.utc)), + ) async def revoke_api_key(self, key_id: str) -> bool: """Revoke an API key.""" - async with self.get_db() as session, session.begin(): - query = ( - update(ApiKey) - .where(col(ApiKey.key_id) == key_id) - .values(revoked_at=datetime.now(timezone.utc)) - ) - result = await session.execute(query) - return getattr(result, "rowcount", 0) > 0 + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + query = ( + update(ApiKey) + .where(col(ApiKey.key_id) == key_id) + .values(revoked_at=datetime.now(timezone.utc)) + ) + result = T.cast(CursorResult, await session.execute(query)) + return result.rowcount > 0 async def delete_api_key(self, key_id: str) -> bool: """Delete an API key.""" - async with self.get_db() as session, session.begin(): - result = await session.execute( - delete(ApiKey).where(col(ApiKey.key_id) == key_id), - ) - return getattr(result, "rowcount", 0) > 0 + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + result = T.cast( + CursorResult, + await session.execute( + delete(ApiKey).where(col(ApiKey.key_id) == key_id) + ), + ) + return result.rowcount > 0 async def insert_persona( self, @@ -932,25 +926,28 @@ class SQLiteDatabase(BaseDatabase): sort_order=0, ): """Insert a new persona record.""" - async with self.get_db() as session, session.begin(): - new_persona = Persona( - persona_id=persona_id, - system_prompt=system_prompt, - begin_dialogs=begin_dialogs or [], - tools=tools, - skills=skills, - custom_error_message=custom_error_message, - folder_id=folder_id, - sort_order=sort_order, - ) - session.add(new_persona) - await session.flush() - await session.refresh(new_persona) - return new_persona + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + new_persona = Persona( + persona_id=persona_id, + system_prompt=system_prompt, + begin_dialogs=begin_dialogs or [], + tools=tools, + skills=skills, + custom_error_message=custom_error_message, + folder_id=folder_id, + sort_order=sort_order, + ) + session.add(new_persona) + await session.flush() + await session.refresh(new_persona) + return new_persona async def get_persona_by_id(self, persona_id): """Get a persona by its ID.""" async with self.get_db() as session: + session: AsyncSession query = select(Persona).where(Persona.persona_id == persona_id) result = await session.execute(query) return result.scalar_one_or_none() @@ -958,6 +955,7 @@ class SQLiteDatabase(BaseDatabase): async def get_personas(self): """Get all personas for a specific bot.""" async with self.get_db() as session: + session: AsyncSession query = select(Persona) result = await session.execute(query) return result.scalars().all() @@ -972,31 +970,39 @@ class SQLiteDatabase(BaseDatabase): custom_error_message=NOT_GIVEN, ): """Update a persona's system prompt or begin dialogs.""" - async with self.get_db() as session, session.begin(): - query = update(Persona).where(col(Persona.persona_id) == persona_id) - values = {} - if system_prompt is not None: - values["system_prompt"] = system_prompt - if begin_dialogs is not None: - values["begin_dialogs"] = begin_dialogs - if tools is not NOT_GIVEN: - values["tools"] = tools - if skills is not NOT_GIVEN: - values["skills"] = skills - if custom_error_message is not NOT_GIVEN: - values["custom_error_message"] = custom_error_message - if not values: - return None - query = query.values(**values) - await session.execute(query) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + query = update(Persona).where(col(Persona.persona_id) == persona_id) + values = {} + if system_prompt is not None: + values["system_prompt"] = system_prompt + if begin_dialogs is not None: + values["begin_dialogs"] = begin_dialogs + if tools is not NOT_GIVEN: + values["tools"] = tools + if skills is not NOT_GIVEN: + values["skills"] = skills + if custom_error_message is not NOT_GIVEN: + values["custom_error_message"] = custom_error_message + if not values: + return None + query = query.values(**values) + await session.execute(query) return await self.get_persona_by_id(persona_id) async def delete_persona(self, persona_id) -> None: """Delete a persona by its ID.""" - async with self.get_db() as session, session.begin(): - await session.execute( - delete(Persona).where(col(Persona.persona_id) == persona_id), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + await session.execute( + delete(Persona).where(col(Persona.persona_id) == persona_id), + ) + + # ==== + # Persona Folder Management + # ==== async def insert_persona_folder( self, @@ -1006,38 +1012,41 @@ class SQLiteDatabase(BaseDatabase): sort_order: int = 0, ) -> PersonaFolder: """Insert a new persona folder.""" - async with self.get_db() as session, session.begin(): - new_folder = PersonaFolder( - name=name, - parent_id=parent_id, - description=description, - sort_order=sort_order, - ) - session.add(new_folder) - await session.flush() - await session.refresh(new_folder) - return new_folder + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + new_folder = PersonaFolder( + name=name, + parent_id=parent_id, + description=description, + sort_order=sort_order, + ) + session.add(new_folder) + await session.flush() + await session.refresh(new_folder) + return new_folder async def get_persona_folder_by_id(self, folder_id: str) -> PersonaFolder | None: """Get a persona folder by its folder_id.""" async with self.get_db() as session: + session: AsyncSession query = select(PersonaFolder).where(PersonaFolder.folder_id == folder_id) result = await session.execute(query) return result.scalar_one_or_none() async def get_persona_folders( - self, - parent_id: str | None = None, + self, parent_id: str | None = None ) -> list[PersonaFolder]: """Get all persona folders, optionally filtered by parent_id. Args: parent_id: If None, returns root folders only. If specified, returns children of that folder. - """ async with self.get_db() as session: + session: AsyncSession if parent_id is None: + # Get root folders (parent_id is NULL) query = ( select(PersonaFolder) .where(col(PersonaFolder.parent_id).is_(None)) @@ -1055,9 +1064,9 @@ class SQLiteDatabase(BaseDatabase): async def get_all_persona_folders(self) -> list[PersonaFolder]: """Get all persona folders.""" async with self.get_db() as session: + session: AsyncSession query = select(PersonaFolder).order_by( - col(PersonaFolder.sort_order), - col(PersonaFolder.name), + col(PersonaFolder.sort_order), col(PersonaFolder.name) ) result = await session.execute(query) return list(result.scalars().all()) @@ -1066,28 +1075,30 @@ class SQLiteDatabase(BaseDatabase): self, folder_id: str, name: str | None = None, - parent_id: Any = NOT_GIVEN, - description: Any = NOT_GIVEN, + parent_id: T.Any = NOT_GIVEN, + description: T.Any = NOT_GIVEN, sort_order: int | None = None, ) -> PersonaFolder | None: """Update a persona folder.""" - async with self.get_db() as session, session.begin(): - query = update(PersonaFolder).where( - col(PersonaFolder.folder_id) == folder_id, - ) - values: dict[str, Any] = {} - if name is not None: - values["name"] = name - if parent_id is not NOT_GIVEN: - values["parent_id"] = parent_id - if description is not NOT_GIVEN: - values["description"] = description - if sort_order is not None: - values["sort_order"] = sort_order - if not values: - return None - query = query.values(**values) - await session.execute(query) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + query = update(PersonaFolder).where( + col(PersonaFolder.folder_id) == folder_id + ) + values: dict[str, T.Any] = {} + if name is not None: + values["name"] = name + if parent_id is not NOT_GIVEN: + values["parent_id"] = parent_id + if description is not NOT_GIVEN: + values["description"] = description + if sort_order is not None: + values["sort_order"] = sort_order + if not values: + return None + query = query.values(**values) + await session.execute(query) return await self.get_persona_folder_by_id(folder_id) async def delete_persona_folder(self, folder_id: str) -> None: @@ -1096,43 +1107,46 @@ class SQLiteDatabase(BaseDatabase): Note: This will also set folder_id to NULL for all personas in this folder, moving them to the root directory. """ - async with self.get_db() as session, session.begin(): - await session.execute( - update(Persona) - .where(col(Persona.folder_id) == folder_id) - .values(folder_id=None), - ) - await session.execute( - delete(PersonaFolder).where( - col(PersonaFolder.folder_id) == folder_id, - ), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + # Move personas to root directory + await session.execute( + update(Persona) + .where(col(Persona.folder_id) == folder_id) + .values(folder_id=None) + ) + # Delete the folder + await session.execute( + delete(PersonaFolder).where( + col(PersonaFolder.folder_id) == folder_id + ), + ) async def move_persona_to_folder( - self, - persona_id: str, - folder_id: str | None, + self, persona_id: str, folder_id: str | None ) -> Persona | None: """Move a persona to a folder (or root if folder_id is None).""" - async with self.get_db() as session, session.begin(): - await session.execute( - update(Persona) - .where(col(Persona.persona_id) == persona_id) - .values(folder_id=folder_id), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + await session.execute( + update(Persona) + .where(col(Persona.persona_id) == persona_id) + .values(folder_id=folder_id) + ) return await self.get_persona_by_id(persona_id) async def get_personas_by_folder( - self, - folder_id: str | None = None, + self, folder_id: str | None = None ) -> list[Persona]: """Get all personas in a specific folder. Args: folder_id: If None, returns personas in root directory. - """ async with self.get_db() as session: + session: AsyncSession if folder_id is None: query = ( select(Persona) @@ -1148,7 +1162,10 @@ class SQLiteDatabase(BaseDatabase): result = await session.execute(query) return list(result.scalars().all()) - async def batch_update_sort_order(self, items: list[dict]) -> None: + async def batch_update_sort_order( + self, + items: list[dict], + ) -> None: """Batch update sort_order for personas and/or folders. Args: @@ -1156,55 +1173,62 @@ class SQLiteDatabase(BaseDatabase): - id: The persona_id or folder_id - type: Either "persona" or "folder" - sort_order: The new sort_order value - """ if not items: return - async with self.get_db() as session, session.begin(): - for item in items: - item_id = item.get("id") - item_type = item.get("type") - sort_order = item.get("sort_order") - if item_id is None or item_type is None or sort_order is None: - continue - if item_type == "persona": - await session.execute( - update(Persona) - .where(col(Persona.persona_id) == item_id) - .values(sort_order=sort_order), - ) - elif item_type == "folder": - await session.execute( - update(PersonaFolder) - .where(col(PersonaFolder.folder_id) == item_id) - .values(sort_order=sort_order), - ) + + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + for item in items: + item_id = item.get("id") + item_type = item.get("type") + sort_order = item.get("sort_order") + + if item_id is None or item_type is None or sort_order is None: + continue + + if item_type == "persona": + await session.execute( + update(Persona) + .where(col(Persona.persona_id) == item_id) + .values(sort_order=sort_order) + ) + elif item_type == "folder": + await session.execute( + update(PersonaFolder) + .where(col(PersonaFolder.folder_id) == item_id) + .values(sort_order=sort_order) + ) async def insert_preference_or_update(self, scope, scope_id, key, value): """Insert a new preference record or update if it exists.""" - async with self.get_db() as session, session.begin(): - query = select(Preference).where( - Preference.scope == scope, - Preference.scope_id == scope_id, - Preference.key == key, - ) - result = await session.execute(query) - existing_preference = result.scalar_one_or_none() - if existing_preference: - existing_preference.value = value - else: - new_preference = Preference( - scope=scope, - scope_id=scope_id, - key=key, - value=value, + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + query = select(Preference).where( + Preference.scope == scope, + Preference.scope_id == scope_id, + Preference.key == key, ) - session.add(new_preference) - return existing_preference or new_preference + result = await session.execute(query) + existing_preference = result.scalar_one_or_none() + if existing_preference: + existing_preference.value = value + else: + new_preference = Preference( + scope=scope, + scope_id=scope_id, + key=key, + value=value, + ) + session.add(new_preference) + return existing_preference or new_preference async def get_preference(self, scope, scope_id, key): """Get a preference by key.""" async with self.get_db() as session: + session: AsyncSession query = select(Preference).where( Preference.scope == scope, Preference.scope_id == scope_id, @@ -1216,6 +1240,7 @@ class SQLiteDatabase(BaseDatabase): async def get_preferences(self, scope, scope_id=None, key=None): """Get all preferences for a specific scope ID or key.""" async with self.get_db() as session: + session: AsyncSession query = select(Preference).where(Preference.scope == scope) if scope_id is not None: query = query.where(Preference.scope_id == scope_id) @@ -1227,6 +1252,7 @@ class SQLiteDatabase(BaseDatabase): async def remove_preference(self, scope, scope_id, key) -> None: """Remove a preference by scope ID and key.""" async with self.get_db() as session: + session: AsyncSession async with session.begin(): await session.execute( delete(Preference).where( @@ -1240,6 +1266,7 @@ class SQLiteDatabase(BaseDatabase): async def clear_preferences(self, scope, scope_id) -> None: """Clear all preferences for a specific scope ID.""" async with self.get_db() as session: + session: AsyncSession async with session.begin(): await session.execute( delete(Preference).where( @@ -1249,12 +1276,18 @@ class SQLiteDatabase(BaseDatabase): ) await session.commit() + # ==== + # Command Configuration & Conflict Tracking + # ==== + async def _run_in_tx( self, fn: Callable[[AsyncSession], Awaitable[TxResult]], ) -> TxResult: - async with self.get_db() as session, session.begin(): - return await fn(session) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + return await fn(session) @staticmethod def _apply_updates(model, **updates) -> None: @@ -1322,11 +1355,16 @@ class SQLiteDatabase(BaseDatabase): async def get_command_configs(self) -> list[CommandConfig]: async with self.get_db() as session: + session: AsyncSession result = await session.execute(select(CommandConfig)) return list(result.scalars().all()) - async def get_command_config(self, handler_full_name: str) -> CommandConfig | None: + async def get_command_config( + self, + handler_full_name: str, + ) -> CommandConfig | None: async with self.get_db() as session: + session: AsyncSession return await session.get(CommandConfig, handler_full_name) async def upsert_command_config( @@ -1405,6 +1443,7 @@ class SQLiteDatabase(BaseDatabase): status: str | None = None, ) -> list[CommandConflict]: async with self.get_db() as session: + session: AsyncSession query = select(CommandConflict) if status: query = query.where(CommandConflict.status == status) @@ -1473,11 +1512,16 @@ class SQLiteDatabase(BaseDatabase): await self._run_in_tx(_op) + # ==== + # Deprecated Methods + # ==== + def get_base_stats(self, offset_sec=86400): """Get base statistics within the specified offset in seconds.""" async def _inner(): async with self.get_db() as session: + session: AsyncSession now = datetime.now() start_time = now - timedelta(seconds=offset_sec) result = await session.execute( @@ -1511,6 +1555,7 @@ class SQLiteDatabase(BaseDatabase): async def _inner(): async with self.get_db() as session: + session: AsyncSession result = await session.execute( select(func.sum(PlatformStat.count)).select_from(PlatformStat), ) @@ -1529,8 +1574,10 @@ class SQLiteDatabase(BaseDatabase): return result def get_grouped_base_stats(self, offset_sec=86400): + # group by platform_id async def _inner(): async with self.get_db() as session: + session: AsyncSession now = datetime.now() start_time = now - timedelta(seconds=offset_sec) result = await session.execute( @@ -1561,6 +1608,10 @@ class SQLiteDatabase(BaseDatabase): t.join() return result + # ==== + # Platform Session Management + # ==== + async def create_platform_session( self, creator: str, @@ -1573,25 +1624,28 @@ class SQLiteDatabase(BaseDatabase): kwargs = {} if session_id: kwargs["session_id"] = session_id - async with self.get_db() as session, session.begin(): - new_session = PlatformSession( - creator=creator, - platform_id=platform_id, - display_name=display_name, - is_group=is_group, - **kwargs, - ) - session.add(new_session) - await session.flush() - await session.refresh(new_session) - return new_session + + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + new_session = PlatformSession( + creator=creator, + platform_id=platform_id, + display_name=display_name, + is_group=is_group, + **kwargs, + ) + session.add(new_session) + await session.flush() + await session.refresh(new_session) + return new_session async def get_platform_session_by_id( - self, - session_id: str, + self, session_id: str ) -> PlatformSession | None: """Get a Platform session by its ID.""" async with self.get_db() as session: + session: AsyncSession query = select(PlatformSession).where( PlatformSession.session_id == session_id, ) @@ -1599,15 +1653,16 @@ class SQLiteDatabase(BaseDatabase): return result.scalar_one_or_none() async def get_platform_sessions_by_ids( - self, - session_ids: list[str], + self, session_ids: list[str] ) -> list[PlatformSession]: """Get platform sessions by IDs.""" if not session_ids: return [] + async with self.get_db() as session: + session: AsyncSession query = select(PlatformSession).where( - col(PlatformSession.session_id).in_(session_ids), + col(PlatformSession.session_id).in_(session_ids) ) result = await session.execute(query) return list(result.scalars().all()) @@ -1640,8 +1695,8 @@ class SQLiteDatabase(BaseDatabase): creator: str, platform_id: str | None = None, exclude_project_sessions: bool = False, - ) -> Any: - query: Any = ( + ): + query = ( select( PlatformSession, col(ChatUIProject.project_id), @@ -1659,20 +1714,23 @@ class SQLiteDatabase(BaseDatabase): ) .where(col(PlatformSession.creator) == creator) ) + if platform_id: query = query.where(PlatformSession.platform_id == platform_id) if exclude_project_sessions: query = query.where(col(ChatUIProject.project_id).is_(None)) + return query @staticmethod - def _rows_to_session_dicts(rows: Sequence[Row[tuple]]) -> list[dict]: + def _rows_to_session_dicts(rows: T.Sequence[Row[tuple]]) -> list[dict]: sessions_with_projects = [] for row in rows: platform_session = row[0] project_id = row[1] project_title = row[2] project_emoji = row[3] + session_dict = { "session": platform_session, "project_id": project_id, @@ -1680,6 +1738,7 @@ class SQLiteDatabase(BaseDatabase): "project_emoji": project_emoji, } sessions_with_projects.append(session_dict) + return sessions_with_projects async def get_platform_sessions_by_creator_paginated( @@ -1692,24 +1751,29 @@ class SQLiteDatabase(BaseDatabase): ) -> tuple[list[dict], int]: """Get paginated Platform sessions for a creator with total count.""" async with self.get_db() as session: + session: AsyncSession offset = (page - 1) * page_size + base_query = self._build_platform_sessions_query( creator=creator, platform_id=platform_id, exclude_project_sessions=exclude_project_sessions, ) + total_result = await session.execute( - select(func.count()).select_from(base_query.subquery()), + select(func.count()).select_from(base_query.subquery()) ) total = int(total_result.scalar_one() or 0) + result_query = ( base_query.order_by(desc(PlatformSession.updated_at)) .offset(offset) .limit(page_size) ) result = await session.execute(result_query) + sessions_with_projects = self._rows_to_session_dicts(result.all()) - return (sessions_with_projects, total) + return sessions_with_projects, total async def update_platform_session( self, @@ -1717,24 +1781,33 @@ class SQLiteDatabase(BaseDatabase): display_name: str | None = None, ) -> None: """Update a Platform session's updated_at timestamp and optionally display_name.""" - async with self.get_db() as session, session.begin(): - values: dict[str, Any] = {"updated_at": datetime.now(timezone.utc)} - if display_name is not None: - values["display_name"] = display_name - await session.execute( - update(PlatformSession) - .where(col(PlatformSession.session_id) == session_id) - .values(**values), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + values: dict[str, T.Any] = {"updated_at": datetime.now(timezone.utc)} + if display_name is not None: + values["display_name"] = display_name + + await session.execute( + update(PlatformSession) + .where(col(PlatformSession.session_id) == session_id) + .values(**values), + ) async def delete_platform_session(self, session_id: str) -> None: """Delete a Platform session by its ID.""" - async with self.get_db() as session, session.begin(): - await session.execute( - delete(PlatformSession).where( - col(PlatformSession.session_id) == session_id, - ), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + await session.execute( + delete(PlatformSession).where( + col(PlatformSession.session_id) == session_id, + ), + ) + + # ==== + # ChatUI Project Management + # ==== async def create_chatui_project( self, @@ -1744,21 +1817,24 @@ class SQLiteDatabase(BaseDatabase): description: str | None = None, ) -> ChatUIProject: """Create a new ChatUI project.""" - async with self.get_db() as session, session.begin(): - project = ChatUIProject( - creator=creator, - title=title, - emoji=emoji, - description=description, - ) - session.add(project) - await session.flush() - await session.refresh(project) - return project + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + project = ChatUIProject( + creator=creator, + title=title, + emoji=emoji, + description=description, + ) + session.add(project) + await session.flush() + await session.refresh(project) + return project async def get_chatui_project_by_id(self, project_id: str) -> ChatUIProject | None: """Get a ChatUI project by its ID.""" async with self.get_db() as session: + session: AsyncSession result = await session.execute( select(ChatUIProject).where( col(ChatUIProject.project_id) == project_id, @@ -1774,6 +1850,7 @@ class SQLiteDatabase(BaseDatabase): ) -> list[ChatUIProject]: """Get all ChatUI projects for a specific creator.""" async with self.get_db() as session: + session: AsyncSession offset = (page - 1) * page_size result = await session.execute( select(ChatUIProject) @@ -1792,33 +1869,40 @@ class SQLiteDatabase(BaseDatabase): description: str | None = None, ) -> None: """Update a ChatUI project.""" - async with self.get_db() as session, session.begin(): - values: dict[str, Any] = {"updated_at": datetime.now(timezone.utc)} - if title is not None: - values["title"] = title - if emoji is not None: - values["emoji"] = emoji - if description is not None: - values["description"] = description - await session.execute( - update(ChatUIProject) - .where(col(ChatUIProject.project_id) == project_id) - .values(**values), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + values: dict[str, T.Any] = {"updated_at": datetime.now(timezone.utc)} + if title is not None: + values["title"] = title + if emoji is not None: + values["emoji"] = emoji + if description is not None: + values["description"] = description + + await session.execute( + update(ChatUIProject) + .where(col(ChatUIProject.project_id) == project_id) + .values(**values), + ) async def delete_chatui_project(self, project_id: str) -> None: """Delete a ChatUI project by its ID.""" - async with self.get_db() as session, session.begin(): - await session.execute( - delete(SessionProjectRelation).where( - col(SessionProjectRelation.project_id) == project_id, - ), - ) - await session.execute( - delete(ChatUIProject).where( - col(ChatUIProject.project_id) == project_id, - ), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + # First remove all session relations + await session.execute( + delete(SessionProjectRelation).where( + col(SessionProjectRelation.project_id) == project_id, + ), + ) + # Then delete the project + await session.execute( + delete(ChatUIProject).where( + col(ChatUIProject.project_id) == project_id, + ), + ) async def add_session_to_project( self, @@ -1826,29 +1910,35 @@ class SQLiteDatabase(BaseDatabase): project_id: str, ) -> SessionProjectRelation: """Add a session to a project.""" - async with self.get_db() as session, session.begin(): - await session.execute( - delete(SessionProjectRelation).where( - col(SessionProjectRelation.session_id) == session_id, - ), - ) - relation = SessionProjectRelation( - session_id=session_id, - project_id=project_id, - ) - session.add(relation) - await session.flush() - await session.refresh(relation) - return relation + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + # First remove existing relation if any + await session.execute( + delete(SessionProjectRelation).where( + col(SessionProjectRelation.session_id) == session_id, + ), + ) + # Then create new relation + relation = SessionProjectRelation( + session_id=session_id, + project_id=project_id, + ) + session.add(relation) + await session.flush() + await session.refresh(relation) + return relation async def remove_session_from_project(self, session_id: str) -> None: """Remove a session from its project.""" - async with self.get_db() as session, session.begin(): - await session.execute( - delete(SessionProjectRelation).where( - col(SessionProjectRelation.session_id) == session_id, - ), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + await session.execute( + delete(SessionProjectRelation).where( + col(SessionProjectRelation.session_id) == session_id, + ), + ) async def get_project_sessions( self, @@ -1858,6 +1948,7 @@ class SQLiteDatabase(BaseDatabase): ) -> list[PlatformSession]: """Get all sessions in a project.""" async with self.get_db() as session: + session: AsyncSession offset = (page - 1) * page_size result = await session.execute( select(PlatformSession) @@ -1874,12 +1965,11 @@ class SQLiteDatabase(BaseDatabase): return list(result.scalars().all()) async def get_project_by_session( - self, - session_id: str, - creator: str, + self, session_id: str, creator: str ) -> ChatUIProject | None: """Get the project that a session belongs to.""" async with self.get_db() as session: + session: AsyncSession result = await session.execute( select(ChatUIProject) .join( @@ -1894,6 +1984,10 @@ class SQLiteDatabase(BaseDatabase): ) return result.scalar_one_or_none() + # ==== + # Cron Job Management + # ==== + async def create_cron_job( self, name: str, @@ -1909,25 +2003,27 @@ class SQLiteDatabase(BaseDatabase): status: str | None = None, job_id: str | None = None, ) -> CronJob: - async with self.get_db() as session, session.begin(): - job = CronJob( - name=name, - job_type=job_type, - cron_expression=cron_expression, - timezone=timezone, - payload=payload or {}, - description=description, - enabled=enabled, - persistent=persistent, - run_once=run_once, - status=status or "scheduled", - ) - if job_id: - job.job_id = job_id - session.add(job) - await session.flush() - await session.refresh(job) - return job + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + job = CronJob( + name=name, + job_type=job_type, + cron_expression=cron_expression, + timezone=timezone, + payload=payload or {}, + description=description, + enabled=enabled, + persistent=persistent, + run_once=run_once, + status=status or "scheduled", + ) + if job_id: + job.job_id = job_id + session.add(job) + await session.flush() + await session.refresh(job) + return job async def update_cron_job( self, @@ -1946,52 +2042,59 @@ class SQLiteDatabase(BaseDatabase): last_run_at: datetime | None | object = CRON_FIELD_NOT_SET, last_error: str | None | object = CRON_FIELD_NOT_SET, ) -> CronJob | None: - async with self.get_db() as session, session.begin(): - updates: dict = {} - for key, val in { - "name": name, - "cron_expression": cron_expression, - "timezone": timezone, - "payload": payload, - "description": description, - "enabled": enabled, - "persistent": persistent, - "run_once": run_once, - "status": status, - "next_run_time": next_run_time, - "last_run_at": last_run_at, - "last_error": last_error, - }.items(): - if val is CRON_FIELD_NOT_SET: - continue - updates[key] = val - stmt = ( - update(CronJob) - .where(col(CronJob.job_id) == job_id) - .values(**updates) - .execution_options(synchronize_session="fetch") - ) - await session.execute(stmt) - result = await session.execute( - select(CronJob).where(col(CronJob.job_id) == job_id), - ) - return result.scalar_one_or_none() + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + updates: dict = {} + for key, val in { + "name": name, + "cron_expression": cron_expression, + "timezone": timezone, + "payload": payload, + "description": description, + "enabled": enabled, + "persistent": persistent, + "run_once": run_once, + "status": status, + "next_run_time": next_run_time, + "last_run_at": last_run_at, + "last_error": last_error, + }.items(): + if val is CRON_FIELD_NOT_SET: + continue + updates[key] = val + + stmt = ( + update(CronJob) + .where(col(CronJob.job_id) == job_id) + .values(**updates) + .execution_options(synchronize_session="fetch") + ) + await session.execute(stmt) + result = await session.execute( + select(CronJob).where(col(CronJob.job_id) == job_id) + ) + return result.scalar_one_or_none() async def delete_cron_job(self, job_id: str) -> None: - async with self.get_db() as session, session.begin(): - await session.execute( - delete(CronJob).where(col(CronJob.job_id) == job_id), - ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + await session.execute( + delete(CronJob).where(col(CronJob.job_id) == job_id) + ) async def get_cron_job(self, job_id: str) -> CronJob | None: async with self.get_db() as session: + session: AsyncSession result = await session.execute( - select(CronJob).where(col(CronJob.job_id) == job_id), + select(CronJob).where(col(CronJob.job_id) == job_id) ) return result.scalar_one_or_none() async def list_cron_jobs(self, job_type: str | None = None) -> list[CronJob]: async with self.get_db() as session: + session: AsyncSession query = select(CronJob) if job_type: query = query.where(col(CronJob.job_type) == job_type) diff --git a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py index b82136368..c1d882656 100644 --- a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py +++ b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py @@ -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: diff --git a/astrbot/core/platform/sources/webchat/webchat_adapter.py b/astrbot/core/platform/sources/webchat/webchat_adapter.py index 906b6eaa4..b4d494b34 100644 --- a/astrbot/core/platform/sources/webchat/webchat_adapter.py +++ b/astrbot/core/platform/sources/webchat/webchat_adapter.py @@ -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) diff --git a/astrbot/dashboard/routes/chat.py b/astrbot/dashboard/routes/chat.py index 8b2ee4f60..d7d4777ac 100644 --- a/astrbot/dashboard/routes/chat.py +++ b/astrbot/dashboard/routes/chat.py @@ -5,10 +5,9 @@ import re import uuid from contextlib import asynccontextmanager from copy import deepcopy -from pathlib import Path, PurePosixPath -from typing import Any, cast +from typing import cast -import anyio +from quart import Response as QuartResponse from quart import g, make_response, request, send_file from astrbot.core import logger, sp @@ -23,26 +22,16 @@ from astrbot.core.platform.sources.webchat.message_parts_helper import ( webchat_message_parts_have_content, ) from astrbot.core.platform.sources.webchat.webchat_queue_mgr import webchat_queue_mgr -from astrbot.core.platform_message_history_mgr import PlatformMessageHistoryManager from astrbot.core.utils.active_event_registry import active_event_registry from astrbot.core.utils.astrbot_path import get_astrbot_data_path from astrbot.core.utils.datetime_utils import to_utc_isoformat from .route import Response, Route, RouteContext +# SSE heartbeat message to keep the connection alive during long-running operations SSE_HEARTBEAT = ": heartbeat\n\n" -def _sanitize_upload_filename(filename: str | None) -> str: - if not filename: - return f"{uuid.uuid4()!s}" - normalized = filename.replace("\\", "/") - name = PurePosixPath(normalized).name.replace("\x00", "").strip() - if name in ("", ".", ".."): - return f"{uuid.uuid4()!s}" - return name - - @asynccontextmanager async def track_conversation(convs: dict, conv_id: str): convs[conv_id] = True @@ -56,196 +45,20 @@ async def _poll_webchat_stream_result(back_queue, username: str): try: result = await asyncio.wait_for(back_queue.get(), timeout=1) except asyncio.TimeoutError: - return (None, False) + # Return a sentinel so the caller can send an SSE heartbeat to + # keep the connection alive during long-running operations (e.g. + # context compression with reasoning models). See #6938. + return None, False except asyncio.CancelledError: - logger.debug(f"[WebChat] 用户 {username} 断开聊天长连接。") - return (None, True) + logger.debug(f"[WebChat] 用户 {username} 断开聊天长连接。") + return None, True except Exception as e: logger.error(f"WebChat stream error: {e}") - return (None, False) - return (result, False) - - -def _resolve_path(path: str) -> Path: - return Path(path).resolve(strict=False) - - -def normalize_legacy_reasoning_message_parts( - message_parts: list[dict] | None, - reasoning: str = "", -) -> list[dict]: - parts: list[dict] = [] - for part in message_parts or []: - if not isinstance(part, dict): - continue - copied = dict(part) - if copied.get("type") == "reasoning": - copied = {"type": "think", "think": copied.get("text", "")} - parts.append(copied) - if reasoning and not any(part.get("type") == "think" for part in parts): - parts.insert(0, {"type": "think", "think": reasoning}) - return parts - - -def extract_reasoning_from_message_parts(message_parts: list[dict]) -> str: - reasoning_parts: list[str] = [] - for part in message_parts: - if part.get("type") != "think": - continue - think = part.get("think") - if isinstance(think, str) and think: - reasoning_parts.append(think) - return "".join(reasoning_parts) - - -def collect_plain_text_from_message_parts(message_parts: list[dict]) -> str: - text_parts: list[str] = [] - for part in message_parts: - if part.get("type") != "plain": - continue - text = part.get("text") - if isinstance(text, str) and text: - text_parts.append(text) - return "".join(text_parts) - - -def build_bot_history_content( - message_parts: list[dict], - *, - agent_stats: dict | None = None, - refs: dict | None = None, - include_legacy_reasoning_field: bool = True, -) -> dict[str, Any]: - normalized_parts = normalize_legacy_reasoning_message_parts(message_parts) - content: dict[str, Any] = {"type": "bot", "message": normalized_parts} - reasoning = extract_reasoning_from_message_parts(normalized_parts) - if reasoning and include_legacy_reasoning_field: - # Keep the legacy field for old clients while the canonical structure - # moves to message parts. - content["reasoning"] = reasoning - if agent_stats: - content["agent_stats"] = agent_stats - if refs: - content["refs"] = refs - return content - - -class BotMessageAccumulator: - def __init__(self) -> None: - self.parts: list[dict] = [] - self.pending_text = "" - self.pending_tool_calls: dict[str, dict] = {} - - def has_content(self) -> bool: - return bool(self.parts or self.pending_text or self.pending_tool_calls) - - def add_plain( - self, - result_text: str, - *, - chain_type: str | None, - streaming: bool, - ) -> None: - if chain_type == "tool_call": - self._flush_pending_text() - self._store_tool_call(result_text) - return - - if chain_type == "tool_call_result": - self._flush_pending_text() - self._store_tool_call_result(result_text) - return - - if chain_type == "reasoning": - self._flush_pending_text() - self._append_think_part(result_text) - return - - if streaming: - self.pending_text += result_text - else: - self.pending_text = result_text - - def add_attachment(self, part: dict | None) -> None: - if not part: - return - self._flush_pending_text() - self.parts.append(part) - - def build_message_parts( - self, *, include_pending_tool_calls: bool = False - ) -> list[dict]: - self._flush_pending_text() - if include_pending_tool_calls and self.pending_tool_calls: - for tool_call in self.pending_tool_calls.values(): - self.parts.append({"type": "tool_call", "tool_calls": [tool_call]}) - self.pending_tool_calls = {} - return self.parts - - def plain_text(self) -> str: - return collect_plain_text_from_message_parts(self.build_message_parts()) - - def reasoning_text(self) -> str: - return extract_reasoning_from_message_parts(self.build_message_parts()) - - def _flush_pending_text(self) -> None: - if not self.pending_text: - return - - if self.parts and self.parts[-1].get("type") == "plain": - last_text = self.parts[-1].get("text") - self.parts[-1]["text"] = f"{last_text or ''}{self.pending_text}" - else: - self.parts.append({"type": "plain", "text": self.pending_text}) - self.pending_text = "" - - def _append_think_part(self, text: str) -> None: - if not text: - return - - if self.parts and self.parts[-1].get("type") == "think": - last_text = self.parts[-1].get("think") - self.parts[-1]["think"] = f"{last_text or ''}{text}" - else: - self.parts.append({"type": "think", "think": text}) - - def _store_tool_call(self, result_text: str) -> None: - tool_call = self._parse_json_object(result_text) - if not tool_call: - return - tool_call_id = str(tool_call.get("id") or "") - if not tool_call_id: - return - self.pending_tool_calls[tool_call_id] = tool_call - - def _store_tool_call_result(self, result_text: str) -> None: - tool_result = self._parse_json_object(result_text) - if not tool_result: - return - - tool_call_id = str(tool_result.get("id") or "") - if not tool_call_id: - return - - tool_call = self.pending_tool_calls.pop(tool_call_id, None) or { - "id": tool_call_id - } - tool_call["result"] = tool_result.get("result") - tool_call["finished_ts"] = tool_result.get("ts") - self.parts.append({"type": "tool_call", "tool_calls": [tool_call]}) - - @staticmethod - def _parse_json_object(raw_text: str) -> dict | None: - try: - parsed = json.loads(raw_text) - except json.JSONDecodeError: - return None - return parsed if isinstance(parsed, dict) else None + return None, False + return result, False class ChatRoute(Route): - platform_history_mgr: PlatformMessageHistoryManager - def __init__( self, context: RouteContext, @@ -280,71 +93,77 @@ class ChatRoute(Route): self.attachments_dir = os.path.join(get_astrbot_data_path(), "attachments") self.legacy_img_dir = os.path.join(get_astrbot_data_path(), "webchat", "imgs") os.makedirs(self.attachments_dir, exist_ok=True) + self.supported_imgs = ["jpg", "jpeg", "png", "gif", "webp"] self.conv_mgr = core_lifecycle.conversation_manager - mgr = core_lifecycle.platform_message_history_manager - assert mgr is not None - self.platform_history_mgr = mgr + self.platform_history_mgr = core_lifecycle.platform_message_history_manager self.db = db self.umop_config_router = core_lifecycle.umop_config_router - assert self.umop_config_router + self.running_convs: dict[str, bool] = {} async def get_file(self): filename = request.args.get("filename") if not filename: - return Response().error("Missing key: filename").to_json() + return Response().error("Missing key: filename").__dict__ + try: file_path = os.path.join(self.attachments_dir, os.path.basename(filename)) - resolved_file_path = _resolve_path(file_path) - resolved_base_dir = _resolve_path(self.attachments_dir) - if not await anyio.Path(resolved_file_path).exists(): + real_file_path = os.path.realpath(file_path) + real_imgs_dir = os.path.realpath(self.attachments_dir) + + if not os.path.exists(real_file_path): + # try legacy file_path = os.path.join( - self.legacy_img_dir, - os.path.basename(filename), + self.legacy_img_dir, os.path.basename(filename) ) - if await anyio.Path(file_path).exists(): - resolved_file_path = _resolve_path(file_path) - resolved_base_dir = _resolve_path(self.legacy_img_dir) - try: - resolved_file_path.relative_to(resolved_base_dir) - except ValueError: - return Response().error("Invalid file path").to_json() + if os.path.exists(file_path): + real_file_path = os.path.realpath(file_path) + real_imgs_dir = os.path.realpath(self.legacy_img_dir) + + if not real_file_path.startswith(real_imgs_dir): + return Response().error("Invalid file path").__dict__ + filename_ext = os.path.splitext(filename)[1].lower() if filename_ext == ".wav": - return await send_file(str(resolved_file_path), mimetype="audio/wav") + return await send_file(real_file_path, mimetype="audio/wav") if filename_ext[1:] in self.supported_imgs: - return await send_file(str(resolved_file_path), mimetype="image/jpeg") - return await send_file(str(resolved_file_path)) + return await send_file(real_file_path, mimetype="image/jpeg") + return await send_file(real_file_path) + except (FileNotFoundError, OSError): - return Response().error("File access error").to_json() + return Response().error("File access error").__dict__ async def get_attachment(self): """Get attachment file by attachment_id.""" attachment_id = request.args.get("attachment_id") if not attachment_id: - return Response().error("Missing key: attachment_id").to_json() + return Response().error("Missing key: attachment_id").__dict__ + try: attachment = await self.db.get_attachment_by_id(attachment_id) if not attachment: - return Response().error("Attachment not found").to_json() + return Response().error("Attachment not found").__dict__ + file_path = attachment.path - resolved_file_path = _resolve_path(file_path) - return await send_file( - str(resolved_file_path), - mimetype=attachment.mime_type, - ) + real_file_path = os.path.realpath(file_path) + + return await send_file(real_file_path, mimetype=attachment.mime_type) + except (FileNotFoundError, OSError): - return Response().error("File access error").to_json() + return Response().error("File access error").__dict__ async def post_file(self): """Upload a file and create an attachment record, return attachment_id.""" post_data = await request.files if "file" not in post_data: - return Response().error("Missing key: file").to_json() + return Response().error("Missing key: file").__dict__ + file = post_data["file"] - filename = _sanitize_upload_filename(file.filename) + filename = file.filename or f"{uuid.uuid4()!s}" content_type = file.content_type or "application/octet-stream" + + # 根据 content_type 判断文件类型并添加扩展名 if content_type.startswith("image"): attach_type = "image" elif content_type.startswith("audio"): @@ -354,22 +173,21 @@ class ChatRoute(Route): else: attach_type = "file" - attachments_dir = Path(self.attachments_dir).resolve(strict=False) - file_path = (attachments_dir / filename).resolve(strict=False) - if not file_path.is_relative_to(attachments_dir): - return Response().error("Invalid filename").__dict__ - - await file.save(str(file_path)) + path = os.path.join(self.attachments_dir, filename) + await file.save(path) # 创建 attachment 记录 attachment = await self.db.insert_attachment( - path=str(file_path), + path=path, type=attach_type, mime_type=content_type, ) + if not attachment: - return Response().error("Failed to create attachment").to_json() + return Response().error("Failed to create attachment").__dict__ + filename = os.path.basename(attachment.path) + return ( Response() .ok( @@ -377,13 +195,13 @@ class ChatRoute(Route): "attachment_id": attachment.attachment_id, "filename": filename, "type": attach_type, - }, + } ) - .to_json() + .__dict__ ) async def _build_user_message_parts(self, message: str | list) -> list[dict]: - """构建用户消息的部分列表。""" + """构建用户消息的部分列表。""" return await build_webchat_message_parts( message, get_attachment_by_id=self.db.get_attachment_by_id, @@ -391,11 +209,9 @@ class ChatRoute(Route): ) 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, @@ -405,9 +221,7 @@ class ChatRoute(Route): ) def _extract_web_search_refs( - self, - accumulated_text: str, - accumulated_parts: list, + self, accumulated_text: str, accumulated_parts: list ) -> dict: """从消息中提取 web_search_tavily 的引用 @@ -416,20 +230,26 @@ class ChatRoute(Route): accumulated_parts: 累积的消息部分列表 Returns: - 包含 used 列表的字典,记录被引用的搜索结果 - + 包含 used 列表的字典,记录被引用的搜索结果 """ - supported = ["web_search_tavily", "web_search_bocha"] + supported = [ + "web_search_baidu", + "web_search_tavily", + "web_search_bocha", + "web_search_brave", + ] + # 从 accumulated_parts 中找到所有 web_search_tavily 的工具调用结果 web_search_results = {} tool_call_parts = [ p for p in accumulated_parts if p.get("type") == "tool_call" and p.get("tool_calls") ] + 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: @@ -443,11 +263,16 @@ class ChatRoute(Route): } except (json.JSONDecodeError, KeyError): pass + if not web_search_results: return {} + + # 从文本中提取所有 xxx 标签并去重 ref_indices = { - m.strip() for m in re.findall("(.*?)", accumulated_text) + m.strip() for m in re.findall(r"(.*?)", accumulated_text) } + + # 构建被引用的结果列表 used_refs = [] for ref_index in ref_indices: if ref_index not in web_search_results: @@ -456,6 +281,7 @@ class ChatRoute(Route): if favicon := sp.temporary_cache.get("_ws_favicon", {}).get(payload["url"]): payload["favicon"] = favicon used_refs.append(payload) + return {"used": used_refs} if used_refs else {} def _sanitize_message_content(self, content: dict) -> dict: @@ -518,8 +344,7 @@ class ChatRoute(Route): async def _delete_threads_by_ids(self, thread_ids: list[str], creator: str) -> None: for thread_id in thread_ids: unified_msg_origin = self._build_thread_unified_msg_origin( - creator, - thread_id, + creator, thread_id ) active_event_registry.request_agent_stop_all(unified_msg_origin) await self.conv_mgr.delete_conversations_by_user_id(unified_msg_origin) @@ -534,7 +359,7 @@ class ChatRoute(Route): async def _load_current_conversation_history(self, session) -> tuple[str, list]: unified_msg_origin = self._build_webchat_unified_msg_origin(session) conversation_id = await self.conv_mgr.get_curr_conversation_id( - unified_msg_origin, + unified_msg_origin ) if not conversation_id: return "", [] @@ -553,9 +378,7 @@ class ChatRoute(Route): return conversation_id, history if isinstance(history, list) else [] def _find_checkpoint_index( - self, - history: list[dict], - checkpoint_id: str, + self, history: list[dict], checkpoint_id: str ) -> int | None: for index, message in enumerate(history): if get_checkpoint_id(message) == checkpoint_id: @@ -563,9 +386,7 @@ class ChatRoute(Route): return None def _find_turn_range( - self, - history: list[dict], - checkpoint_id: str, + self, history: list[dict], checkpoint_id: str ) -> tuple[int, int] | None: checkpoint_index = self._find_checkpoint_index(history, checkpoint_id) if checkpoint_index is None: @@ -649,10 +470,7 @@ class ChatRoute(Route): return result def _find_turn_user_index( - self, - history: list[dict], - start: int, - end: int, + self, history: list[dict], start: int, end: int ) -> int | None: for index in range(start, end): message = history[index] @@ -661,10 +479,7 @@ class ChatRoute(Route): return None def _find_turn_final_assistant_index( - self, - history: list[dict], - start: int, - end: int, + self, history: list[dict], start: int, end: int ) -> int | None: for index in range(end - 1, start - 1, -1): message = history[index] @@ -686,9 +501,7 @@ class ChatRoute(Route): return history_list async def _delete_platform_history_after( - self, - session, - message_id: int, + self, session, message_id: int ) -> list[int]: history_list = await self._get_sorted_platform_history(session) should_delete = False @@ -706,18 +519,27 @@ class ChatRoute(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, platform_history_id: str = "webchat", ): """保存 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}) + + 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 record = await self.platform_history_mgr.insert( platform_id=platform_history_id, @@ -731,16 +553,19 @@ class ChatRoute(Route): async def chat(self, post_data: dict | None = None): username = g.get("username", "guest") + if post_data is None: post_data = await request.json if post_data is None: - return Response().error("Missing JSON body").to_json() + return Response().error("Missing JSON body").__dict__ if "message" not in post_data and "files" not in post_data: - return Response().error("Missing key: message or files").to_json() + return Response().error("Missing key: message or files").__dict__ + if "session_id" not in post_data and "conversation_id" not in post_data: return ( - Response().error("Missing key: session_id or conversation_id").to_json() + Response().error("Missing key: session_id or conversation_id").__dict__ ) + message = post_data["message"] session_id = post_data.get("session_id", post_data.get("conversation_id")) selected_provider = post_data.get("selected_provider") @@ -750,15 +575,19 @@ class ChatRoute(Route): thread_selected_text = post_data.get("_thread_selected_text") if not session_id: - return Response().error("session_id is empty").to_json() + return Response().error("session_id is empty").__dict__ + webchat_conv_id = session_id + + # 构建用户消息段(包含 path 用于传递给 adapter) message_parts = await self._build_user_message_parts(message) if not webchat_message_parts_have_content(message_parts): return ( Response() .error("Message content is empty (reply only is not allowed)") - .to_json() + .__dict__ ) + message_id = str(uuid.uuid4()) llm_checkpoint_id = post_data.get("_llm_checkpoint_id") or str(uuid.uuid4()) skip_user_history = bool(post_data.get("_skip_user_history")) @@ -770,61 +599,14 @@ class ChatRoute(Route): async def stream(): client_disconnected = False - message_accumulator = BotMessageAccumulator() + accumulated_parts = [] + accumulated_text = "" + accumulated_reasoning = "" + tool_calls = {} agent_stats = {} refs = {} - - 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"Failed to extract web search refs: {e}", - exc_info=True, - ) - extracted_refs = refs - - saved_record = await self._save_bot_message( - webchat_conv_id, - message_parts_to_save, - agent_stats, - extracted_refs, - llm_checkpoint_id, - platform_history_id, - ) - message_accumulator = BotMessageAccumulator() - agent_stats = {} - refs = {} - return saved_record - - def build_attachment_saved_event(part: dict | None) -> str | None: - if not part or not part.get("attachment_id") or not part.get("type"): - return None - - payload = { - "type": "attachment_saved", - "data": { - "id": part["attachment_id"], - "type": part["type"], - }, - } - return f"data: {json.dumps(payload, ensure_ascii=False)}\n\n" - try: + # Emit session_id first so clients can bind the stream immediately. session_info = { "type": "session_id", "data": None, @@ -837,7 +619,7 @@ class ChatRoute(Route): "data": { "id": saved_user_record.id, "created_at": to_utc_isoformat( - saved_user_record.created_at, + saved_user_record.created_at ), "llm_checkpoint_id": llm_checkpoint_id, }, @@ -847,26 +629,31 @@ class ChatRoute(Route): async with track_conversation(self.running_convs, webchat_conv_id): while True: result, should_break = await _poll_webchat_stream_result( - back_queue, - username, + back_queue, username ) if should_break: client_disconnected = True break if not result: + # Send an SSE comment as keep-alive so the client + # doesn't time out during slow backend ops like + # context compression with reasoning models (#6938). if not client_disconnected: yield SSE_HEARTBEAT continue + if ( "message_id" in result and result["message_id"] != message_id ): logger.warning("webchat stream message_id mismatch") continue + result_text = result["data"] msg_type = result.get("type") streaming = result.get("streaming", False) chain_type = result.get("chain_type") + if chain_type == "agent_stats": stats_info = { "type": "agent_stats", @@ -875,82 +662,114 @@ class ChatRoute(Route): yield f"data: {json.dumps(stats_info, ensure_ascii=False)}\n\n" agent_stats = stats_info["data"] continue + + # 发送 SSE 数据 try: if not client_disconnected: yield f"data: {json.dumps(result, ensure_ascii=False)}\n\n" except Exception as e: if not client_disconnected: logger.debug( - f"[WebChat] 用户 {username} 断开聊天长连接。 {e}", + f"[WebChat] 用户 {username} 断开聊天长连接。 {e}" ) client_disconnected = True + try: if not client_disconnected: await asyncio.sleep(0.05) except asyncio.CancelledError: - logger.debug(f"[WebChat] 用户 {username} 断开聊天长连接。") + logger.debug(f"[WebChat] 用户 {username} 断开聊天长连接。") client_disconnected = True + + # 累积消息部分 if msg_type == "plain": - message_accumulator.add_plain( - result_text, - chain_type=chain_type, - streaming=streaming, - ) + chain_type = result.get("chain_type") + if chain_type == "tool_call": + 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 = "" + elif chain_type == "tool_call_result": + 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) + elif chain_type == "reasoning": + accumulated_reasoning += result_text + elif streaming: + accumulated_text += result_text + else: + accumulated_text = result_text elif msg_type == "image": filename = result_text.replace("[IMAGE]", "") part = await self._create_attachment_from_file( - filename, - "image", + filename, "image" ) - message_accumulator.add_attachment(part) - if attachment_saved_event := build_attachment_saved_event( - part - ): - yield attachment_saved_event + if part: + accumulated_parts.append(part) elif msg_type == "record": filename = result_text.replace("[RECORD]", "") part = await self._create_attachment_from_file( - filename, - "record", + filename, "record" ) - message_accumulator.add_attachment(part) - if attachment_saved_event := build_attachment_saved_event( - part - ): - yield attachment_saved_event + if part: + accumulated_parts.append(part) elif msg_type == "file": + # 格式: [FILE]filename filename = result_text.replace("[FILE]", "") part = await self._create_attachment_from_file( - filename, - "file", + filename, "file" ) - message_accumulator.add_attachment(part) - if attachment_saved_event := build_attachment_saved_event( - part - ): - yield attachment_saved_event - elif msg_type == "video": - filename = result_text.replace("[VIDEO]", "") - part = await self._create_attachment_from_file( - filename, "video" - ) - message_accumulator.add_attachment(part) - if attachment_saved_event := build_attachment_saved_event( - part - ): - yield attachment_saved_event + if part: + accumulated_parts.append(part) - should_save = False + # 消息结束处理 if msg_type == "end": - should_save = message_accumulator.has_content() or bool( - refs or agent_stats - ) - elif (streaming and msg_type == "complete") or not streaming: - if chain_type not in ("tool_call", "tool_call_result"): - should_save = True + break + elif ( + (streaming and msg_type == "complete") or not streaming + # or msg_type == "break" + ): + if ( + chain_type == "tool_call" + or chain_type == "tool_call_result" + ): + continue - if should_save: - saved_record = await flush_pending_bot_message() + # 提取 web_search_tavily 引用 + try: + refs = self._extract_web_search_refs( + accumulated_text, + accumulated_parts, + ) + except Exception as e: + logger.exception( + f"Failed to extract web search refs: {e}", + exc_info=True, + ) + + saved_record = await self._save_bot_message( + webchat_conv_id, + accumulated_text, + accumulated_parts, + accumulated_reasoning, + agent_stats, + refs, + llm_checkpoint_id, + platform_history_id, + ) # 发送保存的消息信息给前端 if saved_record and not client_disconnected: saved_info = { @@ -958,7 +777,7 @@ class ChatRoute(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, }, @@ -967,20 +786,18 @@ class ChatRoute(Route): yield f"data: {json.dumps(saved_info, ensure_ascii=False)}\n\n" except Exception: pass - if msg_type == "end": - break + accumulated_parts = [] + accumulated_text = "" + accumulated_reasoning = "" + # tool_calls = {} + agent_stats = {} + refs = {} except BaseException as e: logger.exception(f"WebChat stream unexpected error: {e}", exc_info=True) finally: - try: - await flush_pending_bot_message() - except Exception as e: - logger.exception( - f"Failed to persist pending webchat message: {e}", - exc_info=True, - ) webchat_queue_mgr.remove_back_queue(message_id) + # 将消息放入会话特定的队列 chat_queue = webchat_queue_mgr.get_or_create_queue(webchat_conv_id) await chat_queue.put( ( @@ -997,6 +814,7 @@ class ChatRoute(Route): }, ), ) + message_parts_for_storage = strip_message_parts_path_fields(message_parts) if not skip_user_history: @@ -1010,7 +828,7 @@ class ChatRoute(Route): ) response = cast( - "QuartResponse", + QuartResponse, await make_response( stream(), { @@ -1021,52 +839,61 @@ class ChatRoute(Route): }, ), ) - response.timeout = None + response.timeout = None # fix SSE auto disconnect issue return response async def stop_session(self): """Stop active agent runs for a session.""" post_data = await request.json if post_data is None: - return Response().error("Missing JSON body").to_json() + return Response().error("Missing JSON body").__dict__ + session_id = post_data.get("session_id") if not session_id: - return Response().error("Missing key: session_id").to_json() + return Response().error("Missing key: session_id").__dict__ + username = g.get("username", "guest") session = await self.db.get_platform_session_by_id(session_id) if not session: - return Response().error(f"Session {session_id} not found").to_json() + return Response().error(f"Session {session_id} not found").__dict__ if session.creator != username: - return Response().error("Permission denied").to_json() + return Response().error("Permission denied").__dict__ + message_type = ( MessageType.GROUP_MESSAGE.value if session.is_group else MessageType.FRIEND_MESSAGE.value ) - umo = f"{session.platform_id}:{message_type}:{session.platform_id}!{username}!{session_id}" + umo = ( + f"{session.platform_id}:{message_type}:" + f"{session.platform_id}!{username}!{session_id}" + ) stopped_count = active_event_registry.request_agent_stop_all(umo) - return Response().ok(data={"stopped_count": stopped_count}).to_json() + + return Response().ok(data={"stopped_count": stopped_count}).__dict__ async def _delete_session_internal(self, session, username: str) -> None: """Delete a single session and all its related data.""" session_id = session.session_id + + # 删除该会话下的所有对话 message_type = "GroupMessage" if session.is_group else "FriendMessage" unified_msg_origin = f"{session.platform_id}:{message_type}:{session.platform_id}!{username}!{session_id}" - conv_mgr = self.conv_mgr - assert conv_mgr is not None - await conv_mgr.delete_conversations_by_user_id(unified_msg_origin) - mgr = self.platform_history_mgr - assert mgr is not None - history_list = await mgr.get( + await self.conv_mgr.delete_conversations_by_user_id(unified_msg_origin) + + # 获取消息历史中的所有附件 ID 并删除附件 + history_list = await self.platform_history_mgr.get( platform_id=session.platform_id, user_id=session_id, page=1, - page_size=100000, + page_size=100000, # 获取足够多的记录 ) attachment_ids = self._extract_attachment_ids(history_list) if attachment_ids: await self._delete_attachments(attachment_ids) - await mgr.delete( + + # 删除消息历史 + await self.platform_history_mgr.delete( platform_id=session.platform_id, user_id=session_id, offset_sec=99999999, @@ -1076,48 +903,56 @@ class ChatRoute(Route): # 删除与会话关联的配置路由 try: - router = self.umop_config_router - assert router is not None - await router.delete_route(unified_msg_origin) + await self.umop_config_router.delete_route(unified_msg_origin) except ValueError as exc: logger.warning( "Failed to delete UMO route %s during session cleanup: %s", unified_msg_origin, exc, ) + + # 清理队列(仅对 webchat) if session.platform_id == "webchat": webchat_queue_mgr.remove_queues(session_id) + + # 删除会话 await self.db.delete_platform_session(session_id) async def delete_webchat_session(self): """Delete a Platform session and all its related data.""" session_id = request.args.get("session_id") if not session_id: - return Response().error("Missing key: session_id").to_json() + return Response().error("Missing key: session_id").__dict__ username = g.get("username", "guest") + session = await self.db.get_platform_session_by_id(session_id) if not session: - return Response().error(f"Session {session_id} not found").to_json() + return Response().error(f"Session {session_id} not found").__dict__ if session.creator != username: - return Response().error("Permission denied").to_json() + return Response().error("Permission denied").__dict__ + await self._delete_session_internal(session, username) - return Response().ok().to_json() + + return Response().ok().__dict__ async def batch_delete_sessions(self): """Batch delete multiple Platform sessions.""" post_data = await request.json if post_data is None: - return Response().error("Missing JSON body").to_json() + return Response().error("Missing JSON body").__dict__ if not isinstance(post_data, dict): - return Response().error("Invalid JSON body: expected object").to_json() + return Response().error("Invalid JSON body: expected object").__dict__ + session_ids = post_data.get("session_ids") if not session_ids or not isinstance(session_ids, list): - return Response().error("Missing or invalid key: session_ids").to_json() + return Response().error("Missing or invalid key: session_ids").__dict__ + username = g.get("username", "guest") sessions = await self.db.get_platform_sessions_by_ids(session_ids) sessions_by_id = {session.session_id: session for session in sessions} deleted_count = 0 failed_items = [] + for sid in session_ids: session = sessions_by_id.get(sid) if not session: @@ -1126,6 +961,7 @@ class ChatRoute(Route): if session.creator != username: failed_items.append({"session_id": sid, "reason": "permission denied"}) continue + try: await self._delete_session_internal(session, username) deleted_count += 1 @@ -1133,6 +969,7 @@ class ChatRoute(Route): except Exception: logger.warning("Failed to delete session %s", sid) failed_items.append({"session_id": sid, "reason": "internal_error"}) + return ( Response() .ok( @@ -1140,9 +977,9 @@ class ChatRoute(Route): "deleted_count": deleted_count, "failed_count": len(failed_items), "failed_items": failed_items, - }, + } ) - .to_json() + .__dict__ ) def _extract_attachment_ids(self, history_list) -> list[str]: @@ -1159,20 +996,22 @@ class ChatRoute(Route): return attachment_ids async def _delete_attachments(self, attachment_ids: list[str]) -> None: - """删除附件(包括数据库记录和磁盘文件)""" + """删除附件(包括数据库记录和磁盘文件)""" try: attachments = await self.db.get_attachments(attachment_ids) for attachment in attachments: - if not await anyio.Path(attachment.path).exists(): + if not os.path.exists(attachment.path): continue try: - await anyio.Path(attachment.path).unlink() + os.remove(attachment.path) except OSError as e: logger.warning( - f"Failed to delete attachment file {attachment.path}: {e}", + f"Failed to delete attachment file {attachment.path}: {e}" ) except Exception as e: logger.warning(f"Failed to get attachments: {e}") + + # 批量删除数据库记录 try: await self.db.delete_attachments(attachment_ids) except Exception as e: @@ -1181,37 +1020,48 @@ class ChatRoute(Route): async def new_session(self): """Create a new Platform session (default: webchat).""" username = g.get("username", "guest") + + # 获取可选的 platform_id 参数,默认为 webchat platform_id = request.args.get("platform_id", "webchat") + + # 创建新会话 session = await self.db.create_platform_session( creator=username, platform_id=platform_id, is_group=0, ) + return ( Response() .ok( data={ "session_id": session.session_id, "platform_id": session.platform_id, - }, + } ) - .to_json() + .__dict__ ) async def get_sessions(self): """Get all Platform sessions for the current user.""" username = g.get("username", "guest") + + # 获取可选的 platform_id 参数 platform_id = request.args.get("platform_id") + sessions, _ = await self.db.get_platform_sessions_by_creator_paginated( creator=username, platform_id=platform_id, page=1, - page_size=100, + page_size=100, # 暂时返回前100个 exclude_project_sessions=True, ) + + # 转换为字典格式 sessions_data = [] for item in sessions: session = item["session"] + sessions_data.append( { "session_id": session.session_id, @@ -1221,30 +1071,35 @@ class ChatRoute(Route): "is_group": session.is_group, "created_at": to_utc_isoformat(session.created_at), "updated_at": to_utc_isoformat(session.updated_at), - }, + } ) - return Response().ok(data=sessions_data).to_json() + + return Response().ok(data=sessions_data).__dict__ async def get_session(self): """Get session information and message history by session_id.""" session_id = request.args.get("session_id") if not session_id: - return Response().error("Missing key: session_id").to_json() + return Response().error("Missing key: session_id").__dict__ + + # 获取会话信息以确定 platform_id session = await self.db.get_platform_session_by_id(session_id) platform_id = session.platform_id if session else "webchat" + + # 获取项目信息(如果会话属于某个项目) username = g.get("username", "guest") project_info = await self.db.get_project_by_session( - session_id=session_id, - creator=username, + session_id=session_id, creator=username ) - mgr = self.platform_history_mgr - assert mgr is not None - history_ls = await mgr.get( + + # Get platform message history using session_id + history_ls = await self.platform_history_mgr.get( platform_id=platform_id, user_id=session_id, page=1, page_size=1000, ) + history_res = [history.model_dump() for history in history_ls] threads = await self.db.get_webchat_threads_by_parent_session( parent_session_id=session_id, @@ -1256,13 +1111,16 @@ class ChatRoute(Route): "threads": [self._serialize_thread(thread) for thread in threads], "is_running": self.running_convs.get(session_id, False), } + + # 如果会话属于项目,添加项目信息 if project_info: response_data["project"] = { "project_id": project_info.project_id, "title": project_info.title, "emoji": project_info.emoji, } - return Response().ok(data=response_data).to_json() + + return Response().ok(data=response_data).__dict__ async def create_thread(self): """Create or reuse a side thread from a selected assistant message.""" @@ -1293,7 +1151,7 @@ class ChatRoute(Route): return Response().error("Permission denied").__dict__ parent_record = await self.db.get_platform_message_history_by_id( - parent_message_id, + parent_message_id ) if ( not parent_record @@ -1322,7 +1180,7 @@ class ChatRoute(Route): return Response().ok(data=self._serialize_thread(existing)).__dict__ conversation_id, history = await self._load_current_conversation_history( - session, + session ) turn_range = self._find_turn_range(history, checkpoint_id) if not conversation_id or not turn_range: @@ -1373,7 +1231,7 @@ class ChatRoute(Route): "thread": self._serialize_thread(thread), "history": [history.model_dump() for history in history_ls], "is_running": self.running_convs.get(thread_id, False), - }, + } ) .__dict__ ) @@ -1404,7 +1262,7 @@ class ChatRoute(Route): "selected_model": post_data.get("selected_model"), "_platform_history_id": "webchat_thread", "_thread_selected_text": thread.selected_text, - }, + } ) async def delete_thread(self): @@ -1490,7 +1348,7 @@ class ChatRoute(Route): ) conversation_id, history = await self._load_current_conversation_history( - session, + session ) turn_range = self._find_turn_range(history, checkpoint_id) if not conversation_id or not turn_range: @@ -1512,8 +1370,7 @@ class ChatRoute(Route): llm_checkpoint_id=new_checkpoint_id, ) deleted_message_ids = await self._delete_platform_history_after( - session, - message_id, + session, message_id ) thread_ids = await self.db.delete_webchat_threads_by_parent_message_ids( session_id, @@ -1534,7 +1391,7 @@ class ChatRoute(Route): "message": updated.model_dump() if updated else None, "needs_regenerate": True, "truncated_after_message": True, - }, + } ) .__dict__ ) @@ -1582,7 +1439,7 @@ class ChatRoute(Route): return Response().error("Message is not linked to LLM history").__dict__ conversation_id, history = await self._load_current_conversation_history( - session, + session ) turn_range = self._find_turn_range(history, checkpoint_id) if not conversation_id or not turn_range: @@ -1652,26 +1509,34 @@ class ChatRoute(Route): "selected_model": post_data.get("selected_model"), "_skip_user_history": True, "_llm_checkpoint_id": new_checkpoint_id, - }, + } ) async def update_session_display_name(self): """Update a Platform session's display name.""" post_data = await request.json + session_id = post_data.get("session_id") display_name = post_data.get("display_name") + if not session_id: - return Response().error("Missing key: session_id").to_json() + return Response().error("Missing key: session_id").__dict__ if display_name is None: - return Response().error("Missing key: display_name").to_json() + return Response().error("Missing key: display_name").__dict__ + username = g.get("username", "guest") + + # 验证会话是否存在且属于当前用户 session = await self.db.get_platform_session_by_id(session_id) if not session: - return Response().error(f"Session {session_id} not found").to_json() + return Response().error(f"Session {session_id} not found").__dict__ if session.creator != username: - return Response().error("Permission denied").to_json() + return Response().error("Permission denied").__dict__ + + # 更新 display_name await self.db.update_platform_session( session_id=session_id, display_name=display_name, ) - return Response().ok().to_json() + + return Response().ok().__dict__ diff --git a/astrbot/dashboard/routes/live_chat.py b/astrbot/dashboard/routes/live_chat.py index 0b88e5469..16c605848 100644 --- a/astrbot/dashboard/routes/live_chat.py +++ b/astrbot/dashboard/routes/live_chat.py @@ -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) diff --git a/dashboard/src/components/chat/Chat.vue b/dashboard/src/components/chat/Chat.vue index cb2a7e45a..ecf61a937 100644 --- a/dashboard/src/components/chat/Chat.vue +++ b/dashboard/src/components/chat/Chat.vue @@ -1,333 +1,561 @@ diff --git a/dashboard/src/components/chat/ChatInput.vue b/dashboard/src/components/chat/ChatInput.vue index b3d4c8877..0b0aeb3d2 100644 --- a/dashboard/src/components/chat/ChatInput.vue +++ b/dashboard/src/components/chat/ChatInput.vue @@ -1,6 +1,7 @@