Files
script/ai/huawei_gateway.py
T

1087 lines
42 KiB
Python
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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
)
# 401 时刷新 Token 重试
if resp.status_code == 401:
logger.warning(f"compact: 第{i+1}批收到 401,刷新Token重试...")
cache.blacklist_current()
new_token = find_token_in_memory()
if new_token:
headers['Authorization'] = f'Bearer {new_token}'
resp = http_session.request(
method='POST', url=target_url, headers=headers,
data=batch_body, allow_redirects=False,
timeout=UPSTREAM_TIMEOUT_MAX, stream=False
)
if resp.status_code == 401:
cache._blacklist.clear()
# 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()