diff --git a/astrbot/core/agent/context/truncator.py b/astrbot/core/agent/context/truncator.py index c61bc588b..962e2ec33 100644 --- a/astrbot/core/agent/context/truncator.py +++ b/astrbot/core/agent/context/truncator.py @@ -12,14 +12,74 @@ class ContextTruncator: and len(message.tool_calls) > 0 ) + @staticmethod + def _split_system_rest( + messages: list[Message], + ) -> tuple[list[Message], list[Message]]: + """Split messages into system messages and the rest. + + Returns: + tuple: (system_messages, non_system_messages) + """ + first_non_system = 0 + for i, msg in enumerate(messages): + if msg.role != "system": + first_non_system = i + break + return messages[:first_non_system], messages[first_non_system:] + + @staticmethod + def _ensure_user_message( + system_messages: list[Message], + truncated: list[Message], + original_messages: list[Message], + ) -> list[Message]: + """Ensure the result always contains a `user` message immediately after + system messages, as required by some LLM APIs. + + Optimization strategy: + - If `truncated` already begins with a `user` message, return it as-is. + - If a `user` message exists later in `truncated`, move that message to + be the first non-system message while preserving the relative order of + the remaining truncated messages (without mutating the original list). + - Otherwise, fall back to the first `user` message from + `original_messages`. + This reduces unnecessary duplication and ensures the required ordering. + """ + if truncated and truncated[0].role == "user": + return system_messages + truncated + + # If a user message exists inside the truncated list, promote it to the front. + index_in_truncated = next( + (i for i, m in enumerate(truncated) if m.role == "user"), None + ) + if index_in_truncated is not None: + # Build a new truncated list that places the found user message first, + # preserving the order of the other messages and avoiding in-place mutation. + user_msg = truncated[index_in_truncated] + new_truncated = [ + user_msg, + *truncated[:index_in_truncated], + *truncated[index_in_truncated + 1 :], + ] + return system_messages + new_truncated + + # Fallback: find the first user message in the original messages. + first_user = next((m for m in original_messages if m.role == "user"), None) + if first_user is None: + # No user messages at all; return system messages + whatever was truncated. + return system_messages + truncated + + return [*system_messages, first_user, *truncated] + def fix_messages(self, messages: list[Message]) -> list[Message]: - """修复消息列表,确保 tool call 和 tool response 的配对关系有效。 + """Fix the message list to ensure the validity of tool call and tool response pairing. - 此方法确保: - 1. 每个 `tool` 消息前面都有一个包含 tool_calls 的 `assistant` 消息 - 2. 每个包含 tool_calls 的 `assistant` 消息后面都有对应的 `tool` 响应 + This method ensures that: + 1. Each `tool` message is preceded by an `assistant` message containing `tool_calls`. + 2. Each `assistant` message containing `tool_calls` is followed by corresponding ` - 这是 OpenAI Chat Completions API 规范的要求(Gemini 对此执行严格检查)。 + This is a requirement of the OpenAI Chat Completions API specification (Gemini enforces this strictly). """ if not messages: return messages @@ -38,24 +98,25 @@ class ContextTruncator: for msg in messages: if msg.role == "tool": - # 只有在有挂起的 assistant(tool_calls) 时才记录 tool 响应 + # Only record tool responses when there is a pending assistant(tool_calls) if pending_assistant is not None: pending_tools.append(msg) - # else: 孤立的 tool 消息,直接忽略 + # Isolated tool messages without a preceding assistant(tool_calls) are ignored continue if self._has_tool_calls(msg): - # 遇到新的 assistant(tool_calls) 前,先处理旧的 pending 链 + # When encountering a new assistant(tool_calls), first process the old pending chain flush_pending_if_valid() pending_assistant = msg continue - # 非 tool,且不含 tool_calls 的消息 - # 先结束任何 pending 链,再正常追加 + # Non-tool messages that do not contain tool_calls will break the pending chain. + # Flush any pending chain first, then append the current message normally. flush_pending_if_valid() fixed_messages.append(msg) - # 结束时处理最后一个 pending 链 + # Flush the last pending chain at the end, + # ensuring that any remaining valid assistant(tool_calls) and its tools are included in the final list. flush_pending_if_valid() return fixed_messages @@ -66,29 +127,23 @@ class ContextTruncator: keep_most_recent_turns: int, drop_turns: int = 1, ) -> list[Message]: - """截断上下文列表,确保不超过最大长度。 - 一个 turn 包含一个 user 消息和一个 assistant 消息。 - 这个方法会保证截断后的上下文列表符合 OpenAI 的上下文格式。 + """ + Turn-based truncation strategy, which drops the oldest turns while keeping the most recent N turns. + A turn consists of a user message and an assistant message. + This method ensures that the truncated context list conforms to OpenAI's context format. Args: - messages: 上下文列表 - keep_most_recent_turns: 保留最近的对话轮数 - drop_turns: 一次性丢弃的对话轮数 + messages: The original list of messages in the context. + keep_most_recent_turns: The number of most recent turns to keep. If set to -1, it means keeping all turns (no truncation). + drop_turns: The number of turns to drop from the beginning. Returns: - 截断后的上下文列表 + The truncated list of messages. """ if keep_most_recent_turns == -1: return messages - first_non_system = 0 - for i, msg in enumerate(messages): - if msg.role != "system": - first_non_system = i - break - - system_messages = messages[:first_non_system] - non_system_messages = messages[first_non_system:] + system_messages, non_system_messages = self._split_system_rest(messages) if len(non_system_messages) // 2 <= keep_most_recent_turns: return messages @@ -99,7 +154,7 @@ class ContextTruncator: else: truncated_contexts = non_system_messages[-num_to_keep * 2 :] - # 找到第一个 role 为 user 的索引,确保上下文格式正确 + # Find the first user message index = next( (i for i, item in enumerate(truncated_contexts) if item.role == "user"), None, @@ -107,8 +162,9 @@ class ContextTruncator: if index is not None and index > 0: truncated_contexts = truncated_contexts[index:] - result = system_messages + truncated_contexts - + result = self._ensure_user_message( + system_messages, truncated_contexts, messages + ) return self.fix_messages(result) def truncate_by_dropping_oldest_turns( @@ -116,53 +172,39 @@ class ContextTruncator: messages: list[Message], drop_turns: int = 1, ) -> list[Message]: - """丢弃最旧的 N 个对话轮次。""" + """Drop the oldest N turns, regardless of the number of turns to keep.""" if drop_turns <= 0: return messages - first_non_system = 0 - for i, msg in enumerate(messages): - if msg.role != "system": - first_non_system = i - break - - system_messages = messages[:first_non_system] - non_system_messages = messages[first_non_system:] + system_messages, non_system_messages = self._split_system_rest(messages) if len(non_system_messages) // 2 <= drop_turns: truncated_non_system = [] else: truncated_non_system = non_system_messages[drop_turns * 2 :] + # Find the first user message index = next( (i for i, item in enumerate(truncated_non_system) if item.role == "user"), None, ) if index is not None: truncated_non_system = truncated_non_system[index:] - elif truncated_non_system: - truncated_non_system = [] - - result = system_messages + truncated_non_system + result = self._ensure_user_message( + system_messages, truncated_non_system, messages + ) return self.fix_messages(result) def truncate_by_halving( self, messages: list[Message], ) -> list[Message]: - """对半砍策略,删除 50% 的消息""" + """Halve the number of messages, keeping the most recent ones.""" if len(messages) <= 2: return messages - first_non_system = 0 - for i, msg in enumerate(messages): - if msg.role != "system": - first_non_system = i - break - - system_messages = messages[:first_non_system] - non_system_messages = messages[first_non_system:] + system_messages, non_system_messages = self._split_system_rest(messages) messages_to_delete = len(non_system_messages) // 2 if messages_to_delete == 0: @@ -170,6 +212,7 @@ class ContextTruncator: truncated_non_system = non_system_messages[messages_to_delete:] + # Find the first user message index = next( (i for i, item in enumerate(truncated_non_system) if item.role == "user"), None, @@ -177,6 +220,7 @@ class ContextTruncator: if index is not None: truncated_non_system = truncated_non_system[index:] - result = system_messages + truncated_non_system - + result = self._ensure_user_message( + system_messages, truncated_non_system, messages + ) return self.fix_messages(result) diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 356658be6..3a3388141 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -1117,6 +1117,20 @@ CONFIG_METADATA_2 = { "api_base": "https://api.anthropic.com/v1", "timeout": 120, "proxy": "", + "custom_headers": {}, + "anth_thinking_config": {"type": "", "budget": 0, "effort": ""}, + }, + "Kimi Coding Plan": { + "id": "kimi-code", + "provider": "kimi-code", + "type": "kimi_code_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://api.kimi.com/coding/", + "timeout": 120, + "proxy": "", + "custom_headers": {"User-Agent": "claude-code/0.1.0"}, "anth_thinking_config": {"type": "", "budget": 0, "effort": ""}, }, "Moonshot": { diff --git a/astrbot/core/pipeline/process_stage/follow_up.py b/astrbot/core/pipeline/process_stage/follow_up.py index 6c1a4fa06..79ec16a85 100644 --- a/astrbot/core/pipeline/process_stage/follow_up.py +++ b/astrbot/core/pipeline/process_stage/follow_up.py @@ -172,6 +172,9 @@ def try_capture_follow_up(event: AstrMessageEvent) -> FollowUpCapture | None: if not active_sender_id or active_sender_id != sender_id: return None + if runner_event.get_extra("agent_stop_requested"): + return None + ticket = runner.follow_up(message_text=_event_follow_up_text(event)) if not ticket: return None diff --git a/astrbot/core/provider/manager.py b/astrbot/core/provider/manager.py index 874b1c4dc..363035b2a 100644 --- a/astrbot/core/provider/manager.py +++ b/astrbot/core/provider/manager.py @@ -380,6 +380,10 @@ class ProviderManager: from .sources.anthropic_source import ( ProviderAnthropic as ProviderAnthropic, ) + case "kimi_code_chat_completion": + from .sources.kimi_code_source import ( + ProviderKimiCode as ProviderKimiCode, + ) case "googlegenai_chat_completion": from .sources.gemini_source import ( ProviderGoogleGenAI as ProviderGoogleGenAI, diff --git a/astrbot/core/provider/sources/anthropic_source.py b/astrbot/core/provider/sources/anthropic_source.py index 5c34a8d8c..0637ca747 100644 --- a/astrbot/core/provider/sources/anthropic_source.py +++ b/astrbot/core/provider/sources/anthropic_source.py @@ -17,7 +17,6 @@ from astrbot.core.provider.entities import LLMResponse, TokenUsage from astrbot.core.provider.func_tool_manager import ToolSet from astrbot.core.utils.io import download_image_by_url from astrbot.core.utils.network_utils import ( - create_proxy_client, is_connection_error, log_connection_failure, ) @@ -30,6 +29,30 @@ from ..register import register_provider_adapter "Anthropic Claude API 提供商适配器", ) class ProviderAnthropic(Provider): + @staticmethod + def _normalize_custom_headers(provider_config: dict) -> dict[str, str] | None: + custom_headers = provider_config.get("custom_headers", {}) + if not isinstance(custom_headers, dict) or not custom_headers: + return None + normalized_headers: dict[str, str] = {} + for key, value in custom_headers.items(): + normalized_headers[str(key)] = str(value) + return normalized_headers or None + + @classmethod + def _resolve_custom_headers( + cls, + provider_config: dict, + *, + required_headers: dict[str, str] | None = None, + ) -> dict[str, str] | None: + merged_headers = cls._normalize_custom_headers(provider_config) or {} + if required_headers: + for header_name, header_value in required_headers.items(): + if not merged_headers.get(header_name, "").strip(): + merged_headers[header_name] = header_value + return merged_headers or None + def __init__( self, provider_config, @@ -47,6 +70,7 @@ class ProviderAnthropic(Provider): if isinstance(self.timeout, str): self.timeout = int(self.timeout) self.thinking_config = provider_config.get("anth_thinking_config", {}) + self.custom_headers = self._resolve_custom_headers(provider_config) if use_api_key: self._init_api_key(provider_config) @@ -67,7 +91,12 @@ class ProviderAnthropic(Provider): def _create_http_client(self, provider_config: dict) -> httpx.AsyncClient | None: """创建带代理的 HTTP 客户端""" proxy = provider_config.get("proxy", "") - return create_proxy_client("Anthropic", proxy) + if proxy: + logger.info(f"[Anthropic] 使用代理: {proxy}") + return httpx.AsyncClient(proxy=proxy, headers=self.custom_headers) + if self.custom_headers: + return httpx.AsyncClient(headers=self.custom_headers) + return None def _apply_thinking_config(self, payloads: dict) -> None: thinking_type = self.thinking_config.get("type", "") diff --git a/astrbot/core/provider/sources/kimi_code_source.py b/astrbot/core/provider/sources/kimi_code_source.py new file mode 100644 index 000000000..02c200271 --- /dev/null +++ b/astrbot/core/provider/sources/kimi_code_source.py @@ -0,0 +1,27 @@ +from ..register import register_provider_adapter +from .anthropic_source import ProviderAnthropic + +KIMI_CODE_API_BASE = "https://api.kimi.com/coding" +KIMI_CODE_DEFAULT_MODEL = "kimi-for-coding" +KIMI_CODE_USER_AGENT = "claude-code/0.1.0" + + +@register_provider_adapter( + "kimi_code_chat_completion", + "Kimi Code Provider Adapter", +) +class ProviderKimiCode(ProviderAnthropic): + def __init__( + self, + provider_config: dict, + provider_settings: dict, + ) -> None: + merged_provider_config = dict(provider_config) + merged_provider_config.setdefault("api_base", KIMI_CODE_API_BASE) + merged_provider_config.setdefault("model", KIMI_CODE_DEFAULT_MODEL) + merged_provider_config["custom_headers"] = self._resolve_custom_headers( + merged_provider_config, + required_headers={"User-Agent": KIMI_CODE_USER_AGENT}, + ) + + super().__init__(merged_provider_config, provider_settings) diff --git a/astrbot/core/provider/sources/openai_embedding_source.py b/astrbot/core/provider/sources/openai_embedding_source.py index 04397b182..2b62d865c 100644 --- a/astrbot/core/provider/sources/openai_embedding_source.py +++ b/astrbot/core/provider/sources/openai_embedding_source.py @@ -19,17 +19,15 @@ class OpenAIEmbeddingProvider(EmbeddingProvider): self.provider_config = provider_config self.provider_settings = provider_settings proxy = provider_config.get("proxy", "") + provider_id = provider_config.get("id", "unknown_id") http_client = None if proxy: - logger.info(f"[OpenAI Embedding] 使用代理: {proxy}") + logger.info(f"[OpenAI Embedding] {provider_id} Using proxy: {proxy}") http_client = httpx.AsyncClient(proxy=proxy) - api_base = provider_config.get("embedding_api_base", "").strip() - if not api_base: - api_base = "https://api.openai.com/v1" - else: - api_base = api_base.removesuffix("/") - if not api_base.endswith("/v1"): - api_base = f"{api_base}/v1" + api_base = provider_config.get( + "embedding_api_base", "https://api.openai.com/v1" + ).strip() + logger.info(f"[OpenAI Embedding] {provider_id} Using API Base: {api_base}") self.client = AsyncOpenAI( api_key=provider_config.get("embedding_api_key"), base_url=api_base, diff --git a/compose-with-shipyard.yml b/compose-with-shipyard.yml index 24ced5a95..7703293fa 100644 --- a/compose-with-shipyard.yml +++ b/compose-with-shipyard.yml @@ -4,7 +4,10 @@ version: '3.8' services: astrbot: - image: soulter/astrbot:latest + build: + context: . + dockerfile: Dockerfile + image: astrbot:kimi-code container_name: astrbot restart: always ports: # mappings description: https://github.com/AstrBotDevs/AstrBot/issues/497 diff --git a/dashboard/src/components/folder/BaseFolderItemSelector.vue b/dashboard/src/components/folder/BaseFolderItemSelector.vue index 0a421b6a1..ca955ea3d 100644 --- a/dashboard/src/components/folder/BaseFolderItemSelector.vue +++ b/dashboard/src/components/folder/BaseFolderItemSelector.vue @@ -1,24 +1,182 @@