From b98010591e50003768d853e5542f8fec86d9f904 Mon Sep 17 00:00:00 2001 From: chaos Date: Wed, 22 Jul 2026 10:26:35 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0=E8=87=AA=E5=8A=A8?= =?UTF-8?q?=E4=B8=8A=E4=B8=8B=E6=96=87=E6=88=AA=E6=96=AD=E5=8A=9F=E8=83=BD?= =?UTF-8?q?=20(=E5=9F=BA=E4=BA=8E=20LiteLLM=20trim=5Fmessages=20=E6=96=B9?= =?UTF-8?q?=E6=A1=88)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 token 估算函数 (_estimate_tokens, _estimate_messages_tokens) - 新增滑动窗口截断 (_trim_messages): 保留 system/tool 消息, 从最旧对话消息开始丢弃 - 新增单条消息中间截断 (_shorten_message_content): 保留头尾, 迭代逼近目标 token 数 - 新增 _maybe_trim_context: 集成到代理逻辑, 请求体超限时自动截断 - 截断目标: 上游 192K token 限制的 73.5% (~141K), 留出响应空间 - 测试验证: 202K token 请求自动截断至 125K, 上游返回 200 --- ai/huawei_gateway.py | 222 +++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 222 insertions(+) diff --git a/ai/huawei_gateway.py b/ai/huawei_gateway.py index 6762ab4..642df36 100755 --- a/ai/huawei_gateway.py +++ b/ai/huawei_gateway.py @@ -43,6 +43,14 @@ UPSTREAM_TIMEOUT_MAX = 300 # 大请求超时 300s RETRY_ON_429 = 2 # 429限流重试次数 RETRY_ON_504 = 1 # 504超时重试次数 +# ================= 上下文自动截断配置 ================= +# 参考 LiteLLM trim_messages 方案: 滑动窗口 + system/tool 保留 + 中间截断 +MAX_CONTEXT_TOKENS = 196608 # 上游实测上限 192K tokens +RESPONSE_BUDGET = 8192 # 预留 8K tokens 给回复 +TRIM_RATIO = 0.75 # 截断到可用空间的 75% +MAX_TRIM_ATTEMPTS = 5 # 单条消息最大截断尝试次数 +ENABLE_AUTO_TRIM = True # 是否启用自动截断 + # ================= 并发配置 ================= POOL_CONNECTIONS = 64 # 连接池大小(须 ≥ Waitress 线程数) POOL_MAXSIZE = 64 # 单主机最大连接数 @@ -274,6 +282,202 @@ adapter = requests.adapters.HTTPAdapter( http_session.mount('https://', adapter) http_session.mount('http://', adapter) + +# ================= 上下文自动截断 (参考 LiteLLM trim_messages) ================= +def _estimate_tokens(text): + """粗略估算文本的 token 数。 + 英文约 4 字符/token,中文约 1.5 字符/token,混合取 ~3 字符/token。 + """ + if not text: + return 0 + # 统计中文字符比例 + chinese_chars = sum(1 for c in text if '\u4e00' <= c <= '\u9fff') + total_chars = len(text) + if total_chars == 0: + return 0 + chinese_ratio = chinese_chars / total_chars + # 中文多的文本 token 密度更高 + chars_per_token = 4.0 - 2.5 * chinese_ratio # 纯英文=4, 纯中文=1.5 + return max(1, int(total_chars / chars_per_token)) + + +def _estimate_message_tokens(msg): + """估算单条消息的 token 数 (含 role 开销 ~4 tokens)""" + content = msg.get('content', '') + if isinstance(content, list): + # 多模态消息: 提取文本部分 + text_parts = [] + for part in content: + if isinstance(part, dict): + if part.get('type') == 'text': + text_parts.append(part.get('text', '')) + elif part.get('type') == 'image_url': + text_parts.append('') # 图片按 ~512 token 估算 + elif isinstance(part, str): + text_parts.append(part) + text = ' '.join(text_parts) + tokens = _estimate_tokens(text) + 512 * sum( + 1 for p in content if isinstance(p, dict) and p.get('type') == 'image_url' + ) + elif isinstance(content, str): + tokens = _estimate_tokens(content) + else: + tokens = _estimate_tokens(str(content)) + # function_call / tool_calls 额外 token + if 'function_call' in msg or 'tool_calls' in msg: + tokens += _estimate_tokens(json.dumps(msg.get('tool_calls', msg.get('function_call', '')))) + return tokens + 4 # role 开销 + + +def _estimate_messages_tokens(messages): + """估算消息列表的总 token 数""" + return sum(_estimate_message_tokens(m) for m in messages) + + +def _shorten_message_content(content, target_tokens): + """从中间截断消息内容,保留头尾。迭代逼近目标 token 数。""" + if not isinstance(content, str): + return content + if not content: + return content + + current_tokens = _estimate_tokens(content) + if current_tokens <= target_tokens: + return content + + marker = "\n...[truncated]...\n" + for _ in range(MAX_TRIM_ATTEMPTS): + current_tokens = _estimate_tokens(content) + if current_tokens <= target_tokens: + break + # 保守比例: 留 90% 空间避免截断标记导致超限 + ratio = (target_tokens * 0.9) / current_tokens + new_length = max(10, int(len(content) * ratio)) + half = new_length // 2 + content = content[:half] + marker + content[-half:] + + return content + + +def _trim_messages(messages, max_tokens): + """滑动窗口截断消息列表。 + + 策略 (参考 LiteLLM): + 1. 分离 system 消息 (始终保留,超限从中间截断) + 2. 分离末尾 tool 消息 (始终保留) + 3. 对话消息从最新向最旧遍历,超限时丢弃最旧 + 4. 单条消息超限时从中间截断 + """ + if not messages: + return messages + + # 分离 system 消息 + system_messages = [m for m in messages if m.get('role') == 'system'] + non_system = [m for m in messages if m.get('role') != 'system'] + + # 分离末尾连续的 tool 消息 + tool_messages = [] + for m in reversed(non_system): + if m.get('role') != 'tool': + break + tool_messages.append(m) + tool_messages.reverse() + conversation = non_system[:len(non_system) - len(tool_messages)] if tool_messages else non_system + + # 计算 system + tool 的 token 开销 + system_tokens = _estimate_messages_tokens(system_messages) + tool_tokens = _estimate_messages_tokens(tool_messages) + overhead = system_tokens + tool_tokens + + # 如果 system 消息本身就超限,截断 system + available = max_tokens - tool_tokens + if available <= 0: + # tool 消息本身就超限了,只能尽力返回 + return system_messages[:1] + tool_messages if system_messages else tool_messages + + if system_tokens > available * 0.5: + # system 消息占太多,从中间截断每条 system 消息 + target_system_tokens = int(available * 0.3) + for m in system_messages: + current = _estimate_message_tokens(m) + if current > target_system_tokens // len(system_messages): + m['content'] = _shorten_message_content( + m.get('content', ''), target_system_tokens // max(1, len(system_messages)) + ) + system_tokens = _estimate_messages_tokens(system_messages) + + # 剩余给对话消息的空间 + conv_budget = max_tokens - system_tokens - tool_tokens + if conv_budget <= 0: + return system_messages + tool_messages + + # 从最新向最旧遍历,滑动窗口 + final_conv = [] + used = 0 + for msg in reversed(conversation): + msg_tokens = _estimate_message_tokens(msg) + if used + msg_tokens <= conv_budget: + final_conv.insert(0, msg) + used += msg_tokens + else: + # 尝试截断这条消息 + remaining = conv_budget - used + if remaining > 50 and 'function_call' not in msg and 'tool_calls' not in msg: + trimmed = dict(msg) + trimmed['content'] = _shorten_message_content(msg.get('content', ''), remaining - 4) + if _estimate_message_tokens(trimmed) <= remaining: + final_conv.insert(0, trimmed) + used += _estimate_message_tokens(trimmed) + # 空间不够,停止加入更旧的消息 + break + + return system_messages + final_conv + tool_messages + + +def _maybe_trim_context(raw_body): + """检查并截断请求体中的消息列表。返回 (new_body, trimmed_info)""" + if not ENABLE_AUTO_TRIM: + return raw_body, None + + try: + body = json.loads(raw_body) + except (json.JSONDecodeError, UnicodeDecodeError): + return raw_body, None + + messages = body.get('messages') + if not messages or not isinstance(messages, list): + return raw_body, None + + total_tokens = _estimate_messages_tokens(messages) + target_limit = int((MAX_CONTEXT_TOKENS - RESPONSE_BUDGET) * TRIM_RATIO) + + if total_tokens <= target_limit: + return raw_body, None # 无需截断 + + original_count = len(messages) + trimmed_messages = _trim_messages(messages, target_limit) + new_tokens = _estimate_messages_tokens(trimmed_messages) + + if len(trimmed_messages) == original_count and new_tokens >= total_tokens: + return raw_body, None # 截断没效果 + + body['messages'] = trimmed_messages + new_body = json.dumps(body, ensure_ascii=False).encode('utf-8') + + info = { + 'original_messages': original_count, + 'trimmed_messages': len(trimmed_messages), + 'original_tokens_est': total_tokens, + 'trimmed_tokens_est': new_tokens, + 'target_limit': target_limit, + } + logger.info( + f"上下文截断: {original_count}→{len(trimmed_messages)} 条消息, " + f"~{total_tokens}→~{new_tokens} tokens (目标≤{target_limit})" + ) + return new_body, info + + # ================= Flask 应用 ================= app = Flask(__name__) @@ -449,6 +653,24 @@ def proxy(subpath): } }, 413 + # ============ 上下文自动截断 ============ + # 在请求体大小校验之后、转发之前,对 messages 做滑动窗口截断 + if 'chat/completions' in subpath or 'messages' in subpath: + raw_body, trim_info = _maybe_trim_context(raw_body) + if trim_info: + body_size = len(raw_body) # 更新截断后的 body 大小 + # 截断后重新检查是否仍超限 + if body_size > UPSTREAM_BODY_LIMIT: + logger.warning(f"截断后请求体仍超限: {body_size // 1024}KB") + return { + "error": { + "message": f"上下文截断后请求体仍过大({body_size // 1024}KB),请减少请求内容长度。", + "type": "invalid_request_error", + "code": "content_too_large", + "param": None + } + }, 413 + # 动态超时:根据请求体大小自动调整 if body_size > 800 * 1024: upstream_timeout = UPSTREAM_TIMEOUT_MAX