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:
1 parent
dcf9b33fe9
commit
b98010591e
1 file changed
+222
@@ -43,6 +43,14 @@ UPSTREAM_TIMEOUT_MAX = 300 # 大请求超时 300s
|
|||||||
RETRY_ON_429 = 2 # 429限流重试次数
|
RETRY_ON_429 = 2 # 429限流重试次数
|
||||||
RETRY_ON_504 = 1 # 504超时重试次数
|
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_CONNECTIONS = 64 # 连接池大小(须 ≥ Waitress 线程数)
|
||||||
POOL_MAXSIZE = 64 # 单主机最大连接数
|
POOL_MAXSIZE = 64 # 单主机最大连接数
|
||||||
@@ -274,6 +282,202 @@ adapter = requests.adapters.HTTPAdapter(
|
|||||||
http_session.mount('https://', adapter)
|
http_session.mount('https://', adapter)
|
||||||
http_session.mount('http://', 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 应用 =================
|
# ================= Flask 应用 =================
|
||||||
app = Flask(__name__)
|
app = Flask(__name__)
|
||||||
|
|
||||||
@@ -449,6 +653,24 @@ def proxy(subpath):
|
|||||||
}
|
}
|
||||||
}, 413
|
}, 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:
|
if body_size > 800 * 1024:
|
||||||
upstream_timeout = UPSTREAM_TIMEOUT_MAX
|
upstream_timeout = UPSTREAM_TIMEOUT_MAX
|
||||||
|
|||||||
Reference in new issue
Block a user