1072 lines
41 KiB
Python
Executable File
1072 lines
41 KiB
Python
Executable File
#!/usr/bin/env python3
|
||
"""
|
||
华为云 Token 动态网关
|
||
- 6小时缓存机制
|
||
- 支持内存扫描自动刷新
|
||
- HUAWEI_TOKEN 环境变量最高优先级
|
||
- Token 持久化到 /etc/huawei-gateway.env
|
||
- SSE 流式转发
|
||
- 连接池 (20 连接)
|
||
- 401 自动重试
|
||
- 兼容生产环境 (Waitress 32线程 / Gunicorn)
|
||
"""
|
||
import os
|
||
import re
|
||
import sys
|
||
import time
|
||
import gzip
|
||
import logging
|
||
import threading
|
||
import traceback
|
||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||
|
||
from flask import Flask, request, Response
|
||
|
||
# 尝试导入 flask-compress,用于响应压缩
|
||
try:
|
||
from flask_compress import Compress
|
||
_compress_available = True
|
||
except ImportError:
|
||
_compress_available = False
|
||
|
||
# 尝试导入 requests,失败则给出明确提示
|
||
try:
|
||
import requests
|
||
except ImportError:
|
||
print("错误:缺少 requests 模块。请运行: pip install requests")
|
||
sys.exit(1)
|
||
|
||
# ================= 配置 =================
|
||
CACHE_TTL = 19800 # 5.5 小时(安全线)
|
||
MAX_WORKERS = 8 # 内存扫描线程数
|
||
MAX_MEM_SEGMENT = 200 * 1024 * 1024 # 单段最大扫描 200MB
|
||
TOKEN_PATTERN = re.compile(b'Bearer ([A-Za-z0-9+/=_-]{100,})')
|
||
TARGET_HOST = 'tokenhub.developer.huaweicloud.com'
|
||
|
||
# ================= 请求体限制配置 =================
|
||
APIG_BODY_LIMIT = 1200 * 1024 # APIG 请求体限制 ~1.2MB (实测边界1260KB)
|
||
COMPACT_THRESHOLD = 350 * 1024 # 自动压缩阈值:超过350KB就压缩(避免上游504超时)
|
||
GZIP_THRESHOLD = 200 * 1024 # 超过 200KB 时启用 gzip 压缩转发
|
||
UPSTREAM_TIMEOUT_MIN = 60 # 小请求超时 60s
|
||
UPSTREAM_TIMEOUT_MAX = 300 # 大请求超时 300s
|
||
RETRY_ON_429 = 2 # 429限流重试次数
|
||
RETRY_ON_504 = 1 # 504超时重试次数(之后再降级compact)
|
||
|
||
# ================= 日志 =================
|
||
logging.basicConfig(
|
||
level=logging.INFO,
|
||
format='%(asctime)s [%(levelname)s] %(message)s',
|
||
handlers=[
|
||
logging.StreamHandler(sys.stdout)
|
||
]
|
||
)
|
||
logger = logging.getLogger('huawei-gateway')
|
||
|
||
# ================= 缓存 =================
|
||
class TokenCache:
|
||
def __init__(self):
|
||
self._token = None
|
||
self._expires_at = 0
|
||
self._lock = threading.RLock()
|
||
self._last_scan = 0
|
||
self._scan_interval = 60 # 扫描间隔最小 60 秒
|
||
self._blacklist = set() # 已失效的 token 指纹
|
||
|
||
def _fingerprint(self, token):
|
||
"""取 token 前 16 + 后 16 字符做指纹"""
|
||
if len(token) <= 32:
|
||
return token
|
||
return token[:16] + token[-16:]
|
||
|
||
def get(self):
|
||
with self._lock:
|
||
now = time.time()
|
||
if self._token and now < self._expires_at:
|
||
return self._token
|
||
return None
|
||
|
||
def set(self, token, ttl=CACHE_TTL):
|
||
with self._lock:
|
||
self._token = token
|
||
self._expires_at = time.time() + ttl
|
||
self._last_scan = time.time()
|
||
self._blacklist.discard(self._fingerprint(token))
|
||
|
||
def blacklist_current(self):
|
||
"""将当前 token 加入黑名单"""
|
||
with self._lock:
|
||
if self._token:
|
||
self._blacklist.add(self._fingerprint(self._token))
|
||
self._token = None
|
||
self._expires_at = 0
|
||
|
||
def is_blacklisted(self, token):
|
||
with self._lock:
|
||
return self._fingerprint(token) in self._blacklist
|
||
|
||
def is_scan_cooldown(self):
|
||
with self._lock:
|
||
return (time.time() - self._last_scan) < self._scan_interval
|
||
|
||
def clear(self):
|
||
with self._lock:
|
||
self._token = None
|
||
self._expires_at = 0
|
||
|
||
|
||
cache = TokenCache()
|
||
|
||
# ================= 内存扫描 =================
|
||
def scan_pid_mem(pid):
|
||
"""扫描单个进程的内存寻找 Token"""
|
||
maps_path = f'/proc/{pid}/maps'
|
||
mem_path = f'/proc/{pid}/mem'
|
||
|
||
if not os.path.exists(maps_path) or not os.path.exists(mem_path):
|
||
return None
|
||
|
||
try:
|
||
with open(maps_path, 'r') as f:
|
||
for line in f:
|
||
parts = line.split()
|
||
if len(parts) < 2:
|
||
continue
|
||
perms = parts[1]
|
||
if 'r' not in perms or 'w' not in perms:
|
||
continue
|
||
|
||
addrs = parts[0].split('-')
|
||
if len(addrs) != 2:
|
||
continue
|
||
|
||
start = int(addrs[0], 16)
|
||
end = int(addrs[1], 16)
|
||
size = end - start
|
||
|
||
if size > MAX_MEM_SEGMENT or size < 1024:
|
||
continue
|
||
|
||
try:
|
||
with open(mem_path, 'rb') as mem:
|
||
mem.seek(start)
|
||
chunk_size = 64 * 1024
|
||
remaining = size
|
||
while remaining > 0:
|
||
to_read = min(chunk_size, remaining)
|
||
data = mem.read(to_read)
|
||
if not data:
|
||
break
|
||
|
||
for match in TOKEN_PATTERN.finditer(data):
|
||
token = match.group(1).decode('ascii', errors='replace')
|
||
if len(token) > 200:
|
||
return token
|
||
|
||
remaining -= len(data)
|
||
except (PermissionError, OSError, ValueError):
|
||
continue
|
||
|
||
except (PermissionError, OSError, ProcessLookupError):
|
||
pass
|
||
return None
|
||
|
||
|
||
def find_token_in_memory():
|
||
"""在所有进程中扫描 Token"""
|
||
# HUAWEI_TOKEN 环境变量优先级最高
|
||
env_token = os.environ.get('HUAWEI_TOKEN', '').strip()
|
||
if env_token and len(env_token) > 200 and not cache.is_blacklisted(env_token):
|
||
cached = cache.get()
|
||
if cached != env_token:
|
||
cache.set(env_token)
|
||
logger.info("Token 从 HUAWEI_TOKEN 环境变量加载")
|
||
return env_token
|
||
|
||
# 从持久化文件加载
|
||
env_file = '/etc/huawei-gateway.env'
|
||
if os.path.isfile(env_file):
|
||
try:
|
||
with open(env_file, 'r') as f:
|
||
for line in f:
|
||
line = line.strip()
|
||
if line.startswith('HUAWEI_TOKEN='):
|
||
file_token = line.split('=', 1)[1].strip().strip('"').strip("'")
|
||
if file_token and len(file_token) > 200 and not cache.is_blacklisted(file_token):
|
||
cached = cache.get()
|
||
if cached != file_token:
|
||
cache.set(file_token)
|
||
logger.info("Token 从持久化文件加载")
|
||
return file_token
|
||
break
|
||
except (OSError, IOError):
|
||
pass
|
||
|
||
if cache.is_scan_cooldown() and cache.get():
|
||
return cache.get()
|
||
|
||
try:
|
||
pids = [pid for pid in os.listdir('/proc') if pid.isdigit()]
|
||
except OSError:
|
||
logger.error("无法访问 /proc 目录")
|
||
return None
|
||
|
||
# 优先扫描常见进程
|
||
priority_pids = []
|
||
other_pids = []
|
||
|
||
for pid in pids:
|
||
try:
|
||
exe_path = os.readlink(f'/proc/{pid}/exe')
|
||
if any(x in exe_path for x in ['python', 'node', 'java', 'chrome', 'electron']):
|
||
priority_pids.append(pid)
|
||
else:
|
||
other_pids.append(pid)
|
||
except (OSError, PermissionError):
|
||
other_pids.append(pid)
|
||
|
||
all_pids = priority_pids + other_pids
|
||
|
||
found_tokens = []
|
||
with ThreadPoolExecutor(max_workers=MAX_WORKERS) as executor:
|
||
futures = {executor.submit(scan_pid_mem, pid): pid for pid in all_pids}
|
||
for future in as_completed(futures):
|
||
try:
|
||
token = future.result(timeout=5)
|
||
if token and not cache.is_blacklisted(token):
|
||
found_tokens.append((token, futures[future]))
|
||
except Exception:
|
||
continue
|
||
|
||
for token, pid in found_tokens:
|
||
cache.set(token)
|
||
logger.info(f"Token 已刷新 (来源 PID: {pid})")
|
||
return token
|
||
|
||
return cache.get()
|
||
|
||
|
||
# ================= HTTP 会话池 =================
|
||
http_session = requests.Session()
|
||
adapter = requests.adapters.HTTPAdapter(
|
||
pool_connections=20,
|
||
pool_maxsize=20,
|
||
max_retries=3
|
||
)
|
||
http_session.mount('https://', adapter)
|
||
http_session.mount('http://', adapter)
|
||
|
||
# ================= Flask 应用 =================
|
||
app = Flask(__name__)
|
||
|
||
# 响应压缩(flask-compress 与 Waitress SSE 存在兼容问题,暂时禁用)
|
||
# 如需启用,需切换到 Gunicorn 或确保客户端不发送 Accept-Encoding
|
||
if False and _compress_available:
|
||
compress = Compress()
|
||
compress.init_app(app)
|
||
app.config['COMPRESS_MIN_SIZE'] = 256
|
||
app.config['COMPRESS_MIMETYPES'] = [
|
||
'text/html', 'text/plain', 'text/css', 'text/xml',
|
||
'application/json', 'application/javascript',
|
||
'application/xml', 'text/event-stream',
|
||
]
|
||
|
||
|
||
# ================= 自动分批压缩(透明,客户端无感知) =================
|
||
def auto_compact(raw_body, real_token, orig_request):
|
||
"""
|
||
当 chat/completions 请求体超过 APIG 限制时,
|
||
自动分批压缩对话历史,返回标准 OpenAI 格式响应。
|
||
对客户端完全透明。
|
||
"""
|
||
import json as _json
|
||
import uuid
|
||
|
||
try:
|
||
payload = _json.loads(raw_body)
|
||
except Exception:
|
||
return {"error": {"message": "无效的JSON请求体", "type": "invalid_request_error"}}, 400
|
||
|
||
model = payload.get('model', 'glm-5.1')
|
||
messages = payload.get('messages', [])
|
||
max_tokens = payload.get('max_tokens', 500)
|
||
stream = payload.get('stream', False)
|
||
body_size = len(raw_body)
|
||
|
||
if not messages:
|
||
return {"error": {"message": "messages 不能为空", "type": "invalid_request_error"}}, 400
|
||
|
||
# 分离 system 消息和对话消息
|
||
system_msgs = [m for m in messages if m.get('role') == 'system']
|
||
convo_msgs = [m for m in messages if m.get('role') != 'system']
|
||
|
||
if not convo_msgs:
|
||
return {"error": {"message": "无对话内容可压缩", "type": "invalid_request_error"}}, 400
|
||
|
||
# 预截断:每条消息内容限制 500 字,大幅减少 compact 请求的 token 数,避免上游 504
|
||
MAX_MSG_CHARS = 500
|
||
truncated_msgs = []
|
||
for m in convo_msgs:
|
||
content = m.get('content', '') or ''
|
||
if len(content) > MAX_MSG_CHARS:
|
||
truncated_msgs.append({**m, 'content': content[:MAX_MSG_CHARS] + '...[截断]'})
|
||
else:
|
||
truncated_msgs.append(m)
|
||
convo_msgs = truncated_msgs
|
||
|
||
# 按批次分割对话:每批控制在安全大小内
|
||
SAFE_BATCH_BYTES = 350 * 1024 # 每批 350KB(避免上游504超时)
|
||
batches = []
|
||
current_batch = []
|
||
current_size = 0
|
||
|
||
for msg in convo_msgs:
|
||
msg_size = len(_json.dumps(msg, ensure_ascii=False).encode('utf-8'))
|
||
if current_size + msg_size > SAFE_BATCH_BYTES and current_batch:
|
||
batches.append(current_batch)
|
||
current_batch = [msg]
|
||
current_size = msg_size
|
||
else:
|
||
current_batch.append(msg)
|
||
current_size += msg_size
|
||
|
||
if current_batch:
|
||
batches.append(current_batch)
|
||
|
||
num_batches = len(batches)
|
||
logger.info(f"auto_compact: {body_size//1024}KB → {num_batches}批, 每批≤{SAFE_BATCH_BYTES//1024}KB")
|
||
|
||
# 逐批压缩
|
||
target_url = f'https://{TARGET_HOST}/v2/chat/completions'
|
||
summaries = []
|
||
|
||
for i, batch in enumerate(batches):
|
||
batch_messages = system_msgs + batch + [{
|
||
"role": "user",
|
||
"content": "请对以上对话内容进行简洁压缩总结,保留所有关键信息、决策和结论,去除冗余和重复。用简洁的条目式格式输出。"
|
||
}]
|
||
|
||
batch_payload = {
|
||
"model": model,
|
||
"messages": batch_messages,
|
||
"max_tokens": max_tokens,
|
||
"stream": False,
|
||
"temperature": 0.3
|
||
}
|
||
|
||
batch_body = _json.dumps(batch_payload, ensure_ascii=False).encode('utf-8')
|
||
headers = {
|
||
'Authorization': f'Bearer {real_token}',
|
||
'Host': TARGET_HOST,
|
||
'Content-Type': 'application/json',
|
||
}
|
||
|
||
try:
|
||
resp = http_session.request(
|
||
method='POST', url=target_url, headers=headers,
|
||
data=batch_body, allow_redirects=False,
|
||
timeout=UPSTREAM_TIMEOUT_MAX, stream=False
|
||
)
|
||
data = resp.json()
|
||
|
||
if resp.status_code == 200 and 'choices' in data:
|
||
summary = data['choices'][0].get('message', {}).get('content', '')
|
||
summaries.append(summary)
|
||
logger.info(f"auto_compact: 第{i+1}/{num_batches}批完成, {len(summary)}字")
|
||
else:
|
||
err_msg = data.get('error', {}).get('message', str(data)[:200])
|
||
logger.warning(f"auto_compact: 第{i+1}批失败: HTTP {resp.status_code} - {err_msg}")
|
||
summaries.append(f"[第{i+1}批压缩失败,原始{len(batch)}条消息]")
|
||
except Exception as e:
|
||
logger.error(f"auto_compact: 第{i+1}批异常: {e}")
|
||
summaries.append(f"[第{i+1}批压缩异常]")
|
||
|
||
time.sleep(1) # 避免触发限流
|
||
|
||
# 合并所有批次的压缩结果
|
||
if len(summaries) == 1:
|
||
final_summary = summaries[0]
|
||
else:
|
||
merge_messages = system_msgs + [{
|
||
"role": "user",
|
||
"content": "以下是分批压缩的对话摘要,请合并为一个连贯的压缩总结,保留所有关键信息:\n\n" +
|
||
"\n\n---\n\n".join(f"第{i+1}批摘要:\n{s}" for i, s in enumerate(summaries))
|
||
}]
|
||
merge_payload = {
|
||
"model": model,
|
||
"messages": merge_messages,
|
||
"max_tokens": max_tokens,
|
||
"stream": False,
|
||
"temperature": 0.3
|
||
}
|
||
try:
|
||
resp = http_session.request(
|
||
method='POST', url=target_url, headers=headers,
|
||
data=_json.dumps(merge_payload, ensure_ascii=False).encode('utf-8'),
|
||
allow_redirects=False, timeout=UPSTREAM_TIMEOUT_MAX
|
||
)
|
||
merge_data = resp.json()
|
||
if resp.status_code == 200 and 'choices' in merge_data:
|
||
final_summary = merge_data['choices'][0].get('message', {}).get('content', '')
|
||
else:
|
||
final_summary = "\n".join(summaries)
|
||
except Exception:
|
||
final_summary = "\n".join(summaries)
|
||
|
||
logger.info(f"auto_compact: 完成, {body_size//1024}KB → {len(final_summary)}字 ({num_batches}批)")
|
||
|
||
# 返回标准 OpenAI 格式(客户端无感知)
|
||
if stream:
|
||
chat_id = str(uuid.uuid4())
|
||
created = int(time.time())
|
||
|
||
def compact_sse():
|
||
first = {
|
||
"id": chat_id, "object": "chat.completion.chunk", "created": created,
|
||
"model": model,
|
||
"choices": [{"index": 0, "delta": {"role": "assistant", "content": ""}, "finish_reason": None}]
|
||
}
|
||
yield f"data: {_json.dumps(first, ensure_ascii=False)}\n\n"
|
||
|
||
content_chunk = {
|
||
"id": chat_id, "object": "chat.completion.chunk", "created": created,
|
||
"model": model,
|
||
"choices": [{"index": 0, "delta": {"content": final_summary}, "finish_reason": None}]
|
||
}
|
||
yield f"data: {_json.dumps(content_chunk, ensure_ascii=False)}\n\n"
|
||
|
||
done_chunk = {
|
||
"id": chat_id, "object": "chat.completion.chunk", "created": created,
|
||
"model": model,
|
||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]
|
||
}
|
||
yield f"data: {_json.dumps(done_chunk, ensure_ascii=False)}\n\n"
|
||
yield "data: [DONE]\n\n"
|
||
|
||
return Response(compact_sse(), status=200, headers={
|
||
'Content-Type': 'text/event-stream;charset=UTF-8',
|
||
'Cache-Control': 'no-cache',
|
||
'X-Auto-Compact': f'batches={num_batches},original_kb={body_size//1024}',
|
||
})
|
||
else:
|
||
return {
|
||
"id": f"chatcmpl-compact-{int(time.time())}",
|
||
"object": "chat.completion",
|
||
"created": int(time.time()),
|
||
"model": model,
|
||
"choices": [{
|
||
"index": 0,
|
||
"message": {"role": "assistant", "content": final_summary},
|
||
"finish_reason": "stop"
|
||
}],
|
||
"usage": {
|
||
"prompt_tokens": body_size // 4,
|
||
"completion_tokens": len(final_summary),
|
||
"total_tokens": body_size // 4 + len(final_summary)
|
||
}
|
||
}
|
||
|
||
|
||
# ================= 全局请求日志(捕获所有请求,包括404) =================
|
||
@app.before_request
|
||
def log_every_request():
|
||
body_preview = ""
|
||
if request.method in ('POST', 'PUT', 'PATCH') and request.content_length and request.content_length < 2048:
|
||
body_preview = request.get_data()[:200].decode('utf-8', errors='replace')
|
||
logger.info(f">>> {request.method} {request.full_path} | body={request.content_length or 0}bytes | from={request.remote_addr} | {body_preview}")
|
||
|
||
|
||
@app.route('/health')
|
||
def health():
|
||
"""健康检查端点"""
|
||
token = cache.get()
|
||
return {
|
||
"status": "healthy",
|
||
"token_cached": token is not None,
|
||
"token_expires_in": max(0, cache._expires_at - time.time()) if hasattr(cache, '_expires_at') else 0,
|
||
"blacklisted": len(cache._blacklist) if hasattr(cache, '_blacklist') else 0
|
||
}, 200
|
||
|
||
|
||
@app.route('/set_token', methods=['POST'])
|
||
def set_token():
|
||
"""手动注入有效 Token"""
|
||
data = request.get_json(force=True, silent=True) if request.is_json else {}
|
||
token = data.get('token', '')
|
||
if not token or len(token) < 100:
|
||
return {"error": "请提供有效的 token"}, 400
|
||
cache.set(token)
|
||
# 持久化到文件以便重启后恢复
|
||
try:
|
||
with open('/etc/huawei-gateway.env', 'w') as f:
|
||
f.write(f'HUAWEI_TOKEN={token}\n')
|
||
except (OSError, IOError):
|
||
pass
|
||
logger.info("Token 已手动注入并持久化")
|
||
return {"status": "ok", "token_fingerprint": cache._fingerprint(token)}, 200
|
||
|
||
|
||
# ================= Compact 分批压缩端点 =================
|
||
@app.route('/v2/compact', methods=['POST', 'OPTIONS'])
|
||
def compact_endpoint():
|
||
"""
|
||
分批上下文压缩端点:
|
||
- 对话历史超长时,自动分批发送给模型压缩
|
||
- 每批独立总结,最后合并为完整压缩上下文
|
||
- 兼容 OpenAI chat completions 请求格式
|
||
"""
|
||
if request.method == 'OPTIONS':
|
||
return Response(status=200, headers={
|
||
'Access-Control-Allow-Origin': '*',
|
||
'Access-Control-Allow-Methods': 'POST, OPTIONS',
|
||
'Access-Control-Allow-Headers': 'Content-Type, Authorization'
|
||
})
|
||
|
||
import json as _json
|
||
|
||
real_token = find_token_in_memory()
|
||
if not real_token:
|
||
return {"error": {"message": "未找到华为云Token", "type": "server_error"}}, 500
|
||
|
||
try:
|
||
payload = request.get_json(force=True)
|
||
except Exception:
|
||
return {"error": {"message": "无效的JSON请求体", "type": "invalid_request_error"}}, 400
|
||
|
||
model = payload.get('model', 'glm-5.1')
|
||
messages = payload.get('messages', [])
|
||
max_tokens = payload.get('max_tokens', 500)
|
||
stream = payload.get('stream', False)
|
||
|
||
if not messages:
|
||
return {"error": {"message": "messages 不能为空", "type": "invalid_request_error"}}, 400
|
||
|
||
# 计算实际请求体大小(使用原始请求体,而非重新序列化)
|
||
raw_body = request.get_data()
|
||
body_size = len(raw_body)
|
||
|
||
# 如果请求体在限制内,直接转发给上游(无需分批)
|
||
if body_size <= APIG_BODY_LIMIT:
|
||
logger.info(f"compact: 请求体{body_size//1024}KB在限制内,直接转发")
|
||
target_url = f'https://{TARGET_HOST}/v2/chat/completions'
|
||
headers = {
|
||
'Authorization': f'Bearer {real_token}',
|
||
'Host': TARGET_HOST,
|
||
'Content-Type': 'application/json',
|
||
}
|
||
# 保留客户端的其他头
|
||
for k, v in request.headers:
|
||
kl = k.lower()
|
||
if kl not in ('host', 'content-length', 'connection', 'accept-encoding',
|
||
'transfer-encoding', 'authorization', 'content-type'):
|
||
headers[k] = v
|
||
|
||
try:
|
||
resp = http_session.request(
|
||
method='POST', url=target_url, headers=headers,
|
||
data=raw_body, allow_redirects=False,
|
||
timeout=UPSTREAM_TIMEOUT_MAX, stream=True
|
||
)
|
||
except Exception as e:
|
||
logger.error(f"compact 转发失败: {e}")
|
||
return {"error": {"message": f"转发失败: {str(e)}", "type": "server_error"}}, 502
|
||
|
||
# 流式转发
|
||
content_type = resp.headers.get('Content-Type', '')
|
||
if 'text/event-stream' in content_type or stream:
|
||
skip_h = {'transfer-encoding', 'content-encoding', 'content-length',
|
||
'connection', 'keep-alive', 'upgrade'}
|
||
resp_headers = [(k, v) for k, v in resp.headers.items() if k.lower() not in skip_h]
|
||
|
||
def sse_stream():
|
||
try:
|
||
for chunk in resp.iter_content(chunk_size=16384):
|
||
if chunk:
|
||
yield chunk
|
||
finally:
|
||
resp.close()
|
||
return Response(sse_stream(), status=resp.status_code,
|
||
headers=resp_headers, direct_passthrough=True)
|
||
else:
|
||
content = resp.content
|
||
resp.close()
|
||
skip_h2 = {'transfer-encoding', 'content-encoding', 'content-length',
|
||
'connection', 'keep-alive', 'upgrade'}
|
||
return Response(content, status=resp.status_code,
|
||
headers=[(k, v) for k, v in resp.headers.items()
|
||
if k.lower() not in skip_h2])
|
||
|
||
# ============ 分批压缩逻辑 ============
|
||
logger.info(f"compact: 请求体{body_size//1024}KB超限,启动分批压缩")
|
||
|
||
# 分离 system 消息和对话消息
|
||
system_msgs = [m for m in messages if m.get('role') == 'system']
|
||
convo_msgs = [m for m in messages if m.get('role') != 'system']
|
||
|
||
if not convo_msgs:
|
||
return {"error": {"message": "无对话内容可压缩", "type": "invalid_request_error"}}, 400
|
||
|
||
# 按批次分割对话:每批控制在安全大小内
|
||
SAFE_BATCH_BYTES = 350 * 1024 # 每批 350KB(避免上游504超时)
|
||
batches = []
|
||
current_batch = []
|
||
current_size = 0
|
||
|
||
for msg in convo_msgs:
|
||
msg_size = len(_json.dumps(msg, ensure_ascii=False).encode('utf-8'))
|
||
if current_size + msg_size > SAFE_BATCH_BYTES and current_batch:
|
||
batches.append(current_batch)
|
||
current_batch = [msg]
|
||
current_size = msg_size
|
||
else:
|
||
current_batch.append(msg)
|
||
current_size += msg_size
|
||
|
||
if current_batch:
|
||
batches.append(current_batch)
|
||
|
||
logger.info(f"compact: 分为{len(batches)}批, 每批~{SAFE_BATCH_BYTES//1024}KB")
|
||
|
||
# 逐批压缩
|
||
target_url = f'https://{TARGET_HOST}/v2/chat/completions'
|
||
summaries = []
|
||
|
||
for i, batch in enumerate(batches):
|
||
# 构建压缩请求
|
||
batch_messages = system_msgs + batch + [{
|
||
"role": "user",
|
||
"content": "请对以上对话内容进行简洁压缩总结,保留所有关键信息、决策和结论,去除冗余和重复。用简洁的条目式格式输出。"
|
||
}]
|
||
|
||
batch_payload = {
|
||
"model": model,
|
||
"messages": batch_messages,
|
||
"max_tokens": max_tokens,
|
||
"stream": False,
|
||
"temperature": 0.3
|
||
}
|
||
|
||
batch_body = _json.dumps(batch_payload, ensure_ascii=False).encode('utf-8')
|
||
|
||
headers = {
|
||
'Authorization': f'Bearer {real_token}',
|
||
'Host': TARGET_HOST,
|
||
'Content-Type': 'application/json',
|
||
}
|
||
|
||
try:
|
||
resp = http_session.request(
|
||
method='POST', url=target_url, headers=headers,
|
||
data=batch_body, allow_redirects=False,
|
||
timeout=UPSTREAM_TIMEOUT_MAX, stream=False
|
||
)
|
||
|
||
# compact 内部 429/504 重试
|
||
for _retry in range(3):
|
||
if resp.status_code not in (429, 504):
|
||
break
|
||
retry_wait = (_retry + 1) * 5
|
||
logger.warning(f"compact: 第{i+1}批收到 {resp.status_code},等待{retry_wait}s重试...")
|
||
time.sleep(retry_wait)
|
||
resp = http_session.request(
|
||
method='POST', url=target_url, headers=headers,
|
||
data=batch_body, allow_redirects=False,
|
||
timeout=UPSTREAM_TIMEOUT_MAX, stream=False
|
||
)
|
||
|
||
data = resp.json()
|
||
|
||
if resp.status_code == 200 and 'choices' in data:
|
||
summary = data['choices'][0].get('message', {}).get('content', '')
|
||
summaries.append(summary)
|
||
logger.info(f"compact: 第{i+1}/{len(batches)}批压缩完成, {len(summary)}字")
|
||
else:
|
||
err_msg = data.get('error', {}).get('message', str(data)[:200])
|
||
logger.warning(f"compact: 第{i+1}批失败: HTTP {resp.status_code} - {err_msg}")
|
||
# 失败的批次保留原始内容摘要
|
||
fallback = f"[第{i+1}批压缩失败,原始{len(batch)}条消息]"
|
||
summaries.append(fallback)
|
||
except Exception as e:
|
||
logger.error(f"compact: 第{i+1}批异常: {e}")
|
||
summaries.append(f"[第{i+1}批压缩异常]")
|
||
|
||
time.sleep(1) # 避免触发限流
|
||
|
||
# 合并所有批次的压缩结果
|
||
if len(summaries) == 1:
|
||
final_summary = summaries[0]
|
||
else:
|
||
# 对多个摘要做最终合并压缩
|
||
merge_messages = system_msgs + [{
|
||
"role": "user",
|
||
"content": "以下是分批压缩的对话摘要,请合并为一个连贯的压缩总结,保留所有关键信息:\n\n" +
|
||
"\n\n---\n\n".join(f"第{i+1}批摘要:\n{s}" for i, s in enumerate(summaries))
|
||
}]
|
||
|
||
merge_payload = {
|
||
"model": model,
|
||
"messages": merge_messages,
|
||
"max_tokens": max_tokens,
|
||
"stream": False,
|
||
"temperature": 0.3
|
||
}
|
||
|
||
try:
|
||
resp = http_session.request(
|
||
method='POST', url=target_url, headers=headers,
|
||
data=_json.dumps(merge_payload, ensure_ascii=False).encode('utf-8'),
|
||
allow_redirects=False, timeout=UPSTREAM_TIMEOUT_MAX
|
||
)
|
||
merge_data = resp.json()
|
||
if resp.status_code == 200 and 'choices' in merge_data:
|
||
final_summary = merge_data['choices'][0].get('message', {}).get('content', '')
|
||
else:
|
||
final_summary = "\n".join(summaries)
|
||
except Exception:
|
||
final_summary = "\n".join(summaries)
|
||
|
||
logger.info(f"compact: 压缩完成, 最终{len(final_summary)}字 (原始{body_size//1024}KB)")
|
||
|
||
# 返回 OpenAI 兼容格式
|
||
if stream:
|
||
# 流式返回:将压缩结果包装为 SSE 事件
|
||
import uuid
|
||
chat_id = str(uuid.uuid4())
|
||
created = int(time.time())
|
||
|
||
def compact_sse():
|
||
# 首个 chunk:role
|
||
first = {
|
||
"id": chat_id, "object": "chat.completion.chunk", "created": created,
|
||
"model": model,
|
||
"choices": [{"index": 0, "delta": {"role": "assistant", "content": ""}, "finish_reason": None}]
|
||
}
|
||
yield f"data: {_json.dumps(first, ensure_ascii=False)}\n\n"
|
||
|
||
# 内容 chunk
|
||
content_chunk = {
|
||
"id": chat_id, "object": "chat.completion.chunk", "created": created,
|
||
"model": model,
|
||
"choices": [{"index": 0, "delta": {"content": final_summary}, "finish_reason": None}]
|
||
}
|
||
yield f"data: {_json.dumps(content_chunk, ensure_ascii=False)}\n\n"
|
||
|
||
# 结束 chunk
|
||
done_chunk = {
|
||
"id": chat_id, "object": "chat.completion.chunk", "created": created,
|
||
"model": model,
|
||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]
|
||
}
|
||
yield f"data: {_json.dumps(done_chunk, ensure_ascii=False)}\n\n"
|
||
yield "data: [DONE]\n\n"
|
||
|
||
return Response(compact_sse(), status=200, headers={
|
||
'Content-Type': 'text/event-stream;charset=UTF-8',
|
||
'Cache-Control': 'no-cache',
|
||
'X-Compact-Batches': str(len(batches)),
|
||
'X-Compact-Original-KB': str(body_size // 1024),
|
||
})
|
||
else:
|
||
# 非流式返回
|
||
return {
|
||
"id": f"compact-{int(time.time())}",
|
||
"object": "chat.completion",
|
||
"created": int(time.time()),
|
||
"model": model,
|
||
"choices": [{
|
||
"index": 0,
|
||
"message": {"role": "assistant", "content": final_summary},
|
||
"finish_reason": "stop"
|
||
}],
|
||
"usage": {
|
||
"prompt_tokens": body_size // 4, # 估算
|
||
"completion_tokens": len(final_summary),
|
||
"total_tokens": body_size // 4 + len(final_summary)
|
||
},
|
||
"compact_meta": {
|
||
"batches": len(batches),
|
||
"original_kb": body_size // 1024,
|
||
}
|
||
}
|
||
|
||
|
||
@app.route('/v1/<path:subpath>', methods=['POST', 'GET', 'OPTIONS', 'PUT', 'DELETE'])
|
||
@app.route('/v2/<path:subpath>', methods=['POST', 'GET', 'OPTIONS', 'PUT', 'DELETE'])
|
||
def proxy(subpath):
|
||
if request.method == 'OPTIONS':
|
||
return Response(status=200, headers={
|
||
'Access-Control-Allow-Origin': '*',
|
||
'Access-Control-Allow-Methods': 'GET, POST, PUT, DELETE, OPTIONS',
|
||
'Access-Control-Allow-Headers': 'Content-Type, Authorization'
|
||
})
|
||
|
||
real_token = find_token_in_memory()
|
||
if not real_token:
|
||
logger.error("未在内存中找到华为云 Token")
|
||
return {"error": "未在内存中找到华为云Token,请确保华为云相关应用正在运行"}, 500
|
||
|
||
# 构建请求头
|
||
headers = {}
|
||
for k, v in request.headers:
|
||
kl = k.lower()
|
||
if kl not in ('host', 'content-length', 'connection', 'accept-encoding', 'transfer-encoding'):
|
||
headers[k] = v
|
||
|
||
headers['Authorization'] = f'Bearer {real_token}'
|
||
headers['Host'] = TARGET_HOST
|
||
|
||
target_url = f'https://{TARGET_HOST}/v2/{subpath}'
|
||
|
||
# ============ /models 请求直接返回,不转发上游 ============
|
||
if subpath in ('models', 'models/'):
|
||
return {
|
||
"object": "list",
|
||
"data": [
|
||
{"id": "glm-5.1", "object": "model", "owned": "zhipu"},
|
||
]
|
||
}
|
||
|
||
# ============ 请求体处理:大小校验 + 动态超时 ============
|
||
raw_body = request.get_data()
|
||
body_size = len(raw_body)
|
||
|
||
# 动态超时:根据请求体大小自动调整
|
||
if body_size > 800 * 1024:
|
||
upstream_timeout = UPSTREAM_TIMEOUT_MAX
|
||
elif body_size > 200 * 1024:
|
||
upstream_timeout = 180
|
||
else:
|
||
upstream_timeout = UPSTREAM_TIMEOUT_MIN
|
||
|
||
# ============ 自动压缩(仅对 chat/completions) ============
|
||
# 1. 超过 APIG 限制 → 必须压缩(否则请求会被截断)
|
||
# 2. 超过 COMPACT_THRESHOLD → 主动压缩(避免上游504超时)
|
||
if body_size > COMPACT_THRESHOLD and subpath in ('chat/completions', 'chat/completions/'):
|
||
reason = "超APIG限制" if body_size > APIG_BODY_LIMIT else "可能超时"
|
||
logger.info(f"chat/completions 请求体 {body_size//1024}KB > {COMPACT_THRESHOLD//1024}KB ({reason}), 自动触发分批压缩")
|
||
return auto_compact(raw_body, real_token, request)
|
||
|
||
# 非 chat/completions 超限,返回清晰错误
|
||
if body_size > APIG_BODY_LIMIT:
|
||
logger.warning(f"非chat请求体超限: {body_size//1024}KB > {APIG_BODY_LIMIT//1024}KB")
|
||
return {
|
||
"error": {
|
||
"message": f"请求体过大({body_size//1024}KB),超过API网关限制({APIG_BODY_LIMIT//1024}KB)。请减少请求内容长度。",
|
||
"type": "invalid_request_error",
|
||
"code": "content_too_large",
|
||
"param": None
|
||
}
|
||
}, 413
|
||
|
||
try:
|
||
# 使用 stream=True 支持 SSE 流式转发
|
||
resp = http_session.request(
|
||
method=request.method,
|
||
url=target_url,
|
||
headers=headers,
|
||
data=raw_body,
|
||
cookies=request.cookies,
|
||
allow_redirects=False,
|
||
timeout=upstream_timeout,
|
||
stream=True
|
||
)
|
||
|
||
# 401 兜底:Token 可能提前过期,加入黑名单后强制刷新重试
|
||
if resp.status_code == 401:
|
||
resp.close()
|
||
logger.warning("收到 401,将当前 Token 加入黑名单并强制刷新...")
|
||
cache.blacklist_current()
|
||
new_token = find_token_in_memory()
|
||
if new_token and new_token != real_token:
|
||
headers['Authorization'] = f'Bearer {new_token}'
|
||
resp = http_session.request(
|
||
method=request.method,
|
||
url=target_url,
|
||
headers=headers,
|
||
data=raw_body,
|
||
cookies=request.cookies,
|
||
allow_redirects=False,
|
||
timeout=upstream_timeout,
|
||
stream=True
|
||
)
|
||
# 如果新 Token 也 401,清空黑名单避免锁死
|
||
if resp.status_code == 401:
|
||
resp.close()
|
||
logger.warning("新 Token 也 401,清空黑名单避免锁死")
|
||
cache._blacklist.clear()
|
||
# 再试一次
|
||
new_token2 = find_token_in_memory()
|
||
if new_token2:
|
||
headers['Authorization'] = f'Bearer {new_token2}'
|
||
resp = http_session.request(
|
||
method=request.method,
|
||
url=target_url,
|
||
headers=headers,
|
||
data=raw_body,
|
||
cookies=request.cookies,
|
||
allow_redirects=False,
|
||
timeout=upstream_timeout,
|
||
stream=True
|
||
)
|
||
|
||
# ============ 429 限流重试 ============
|
||
if resp.status_code == 429:
|
||
for retry_i in range(1, RETRY_ON_429 + 1):
|
||
resp.close()
|
||
wait = retry_i * 5 # 5s, 10s
|
||
logger.warning(f"上游 429 限流,等待{wait}s后重试({retry_i}/{RETRY_ON_429})...")
|
||
time.sleep(wait)
|
||
resp = http_session.request(
|
||
method=request.method, url=target_url, headers=headers,
|
||
data=raw_body, cookies=request.cookies,
|
||
allow_redirects=False, timeout=upstream_timeout, stream=True
|
||
)
|
||
if resp.status_code != 429:
|
||
break
|
||
logger.warning(f"重试仍返回 429")
|
||
|
||
# ============ 504 超时重试 + 降级compact ============
|
||
if resp.status_code == 504 and subpath in ('chat/completions', 'chat/completions/'):
|
||
for retry_i in range(1, RETRY_ON_504 + 1):
|
||
resp.close()
|
||
wait = retry_i * 3
|
||
logger.warning(f"上游 504 超时,等待{wait}s后重试({retry_i}/{RETRY_ON_504})...")
|
||
time.sleep(wait)
|
||
resp = http_session.request(
|
||
method=request.method, url=target_url, headers=headers,
|
||
data=raw_body, cookies=request.cookies,
|
||
allow_redirects=False, timeout=upstream_timeout, stream=True
|
||
)
|
||
if resp.status_code != 504:
|
||
break
|
||
|
||
# 重试仍504 → 降级为compact压缩后重试
|
||
if resp.status_code == 504:
|
||
resp.close()
|
||
logger.warning(f"504重试仍失败,降级为auto_compact压缩后重试...")
|
||
try:
|
||
return auto_compact(raw_body, real_token, request)
|
||
except Exception as e:
|
||
logger.error(f"降级compact也失败: {e}")
|
||
# compact也失败,返回友好错误
|
||
return {
|
||
"error": {
|
||
"message": "模型推理超时,已尝试压缩上下文但仍失败。请缩短对话后重试。",
|
||
"type": "server_error",
|
||
"code": "model_timeout",
|
||
"param": None
|
||
}
|
||
}, 504
|
||
|
||
# 过滤 hop-by-hop 头和压缩编码头
|
||
skip_headers = {'transfer-encoding', 'content-encoding', 'content-length',
|
||
'connection', 'keep-alive', 'upgrade'}
|
||
response_headers = []
|
||
for k, v in resp.headers.items():
|
||
if k.lower() not in skip_headers:
|
||
response_headers.append((k, v))
|
||
|
||
# 记录上游非200响应
|
||
if resp.status_code != 200:
|
||
try:
|
||
_err_peek = resp.content[:500]
|
||
logger.warning(f"上游返回 {resp.status_code}: {_err_peek.decode('utf-8', errors='replace')}")
|
||
except Exception:
|
||
logger.warning(f"上游返回 {resp.status_code}")
|
||
|
||
# SSE 流式转发
|
||
content_type = resp.headers.get('Content-Type', '')
|
||
if 'text/event-stream' in content_type or resp.headers.get('Transfer-Encoding', '') == 'chunked':
|
||
def sse_stream():
|
||
try:
|
||
for chunk in resp.iter_content(chunk_size=16384):
|
||
if chunk:
|
||
yield chunk
|
||
finally:
|
||
resp.close()
|
||
|
||
return Response(
|
||
sse_stream(),
|
||
status=resp.status_code,
|
||
headers=response_headers,
|
||
direct_passthrough=True
|
||
)
|
||
else:
|
||
# 非流式响应:读取完整内容
|
||
content = resp.content
|
||
resp.close()
|
||
|
||
# 华为云 ModelArts 错误 → OpenAI 标准格式
|
||
if resp.status_code >= 400:
|
||
try:
|
||
import json as _jj
|
||
err = _jj.loads(content)
|
||
if 'error_code' in err and 'error_msg' in err:
|
||
openai_err = {
|
||
"error": {
|
||
"message": err.get('error_msg', ''),
|
||
"type": err.get('error', {}).get('type', 'server_error'),
|
||
"code": err.get('error_code', ''),
|
||
"param": None
|
||
}
|
||
}
|
||
content = _jj.dumps(openai_err).encode('utf-8')
|
||
response_headers = [(k, v) for k, v in response_headers if k.lower() != 'content-length']
|
||
except Exception:
|
||
pass
|
||
|
||
return Response(
|
||
content,
|
||
status=resp.status_code,
|
||
headers=response_headers
|
||
)
|
||
|
||
except requests.exceptions.Timeout:
|
||
logger.error("请求华为云 API 超时")
|
||
return {"error": "网关超时,请稍后重试"}, 504
|
||
except requests.exceptions.ConnectionError:
|
||
logger.error("无法连接到华为云 API")
|
||
return {"error": "无法连接到华为云服务"}, 502
|
||
except Exception as e:
|
||
logger.error(f"网关转发失败: {traceback.format_exc()}")
|
||
return {"error": f"网关转发失败: {str(e)}"}, 500
|
||
|
||
|
||
def main():
|
||
port = int(sys.argv[1]) if len(sys.argv) > 1 else 8080
|
||
host = sys.argv[2] if len(sys.argv) > 2 else '127.0.0.1'
|
||
|
||
# 从持久化文件加载 token
|
||
env_file = '/etc/huawei-gateway.env'
|
||
if os.path.isfile(env_file):
|
||
try:
|
||
with open(env_file, 'r') as f:
|
||
for line in f:
|
||
line = line.strip()
|
||
if line.startswith('HUAWEI_TOKEN=') and 'HUAWEI_TOKEN' not in os.environ:
|
||
val = line.split('=', 1)[1].strip().strip('"').strip("'")
|
||
if val and len(val) > 200:
|
||
os.environ['HUAWEI_TOKEN'] = val
|
||
logger.info("从持久化文件恢复 Token")
|
||
break
|
||
except (OSError, IOError):
|
||
pass
|
||
|
||
# 尝试使用生产级 WSGI 服务器
|
||
try:
|
||
import waitress
|
||
logger.info(f"使用 Waitress 启动网关 ({host}:{port})")
|
||
waitress.serve(app, host=host, port=port, threads=32)
|
||
except ImportError:
|
||
try:
|
||
import gunicorn.app.wsgiapp
|
||
logger.info(f"使用 Gunicorn 启动网关 ({host}:{port})")
|
||
os.execlp('gunicorn', 'gunicorn', '-w', '4', '-b', f'{host}:{port}', '--access-logfile', '-', 'huawei_gateway:app')
|
||
except (ImportError, OSError):
|
||
logger.warning("未安装 Waitress/Gunicorn,使用 Flask 开发服务器(建议生产环境安装 waitress)")
|
||
logger.info(f"启动网关 ({host}:{port})")
|
||
app.run(host=host, port=port, debug=False, threaded=True)
|
||
|
||
|
||
if __name__ == '__main__':
|
||
main()
|