feat: 添加自动上下文截断功能 (基于 LiteLLM trim_messages 方案)

- 新增 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
This commit is contained in:
chaos committed 2026-07-22 10:26:35 +08:00
1 parent dcf9b33fe9
commit b98010591e
1 file changed
+222
+222
View File
@@ -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