- anthropic_request_to_openai(): Anthropic 请求转 OpenAI 格式 - system 提取到 messages 首条 - content blocks 展开 (text/tool_use/tool_result) - tools 格式转换 (input_schema → parameters) - max_tokens/temperature/stop_sequences 等参数映射 - openai_response_to_anthropic(): OpenAI 响应转 Anthropic 格式 - content 数组构建 (text/tool_use) - stop_reason 映射 (stop→end_turn, tool_calls→tool_use, length→max_tokens) - usage 字段映射 (input_tokens/output_tokens) - openai_stream_to_anthropic_stream(): 流式 SSE 转换生成器 - message_start/content_block_start/content_block_delta/content_block_stop - message_delta/message_stop 事件序列 - 支持 text_delta 和 input_json_delta - 正确的 block index 分配 - /v1/messages 路由: 复用现有 token 获取/超时/重试逻辑 - 非流式: 请求转换→上游转发→响应转换 - 流式: 请求转换→上游流式转发→SSE 事件转换 测试通过: 非流式/流式文本、system prompt、tool_use、tool_result 多轮对话
1320 lines
50 KiB
Python
Executable File
1320 lines
50 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 json
|
||
import time
|
||
import logging
|
||
import threading
|
||
import traceback
|
||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||
|
||
from flask import Flask, request, Response
|
||
|
||
# 尝试导入 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'
|
||
|
||
# ================= 请求体限制配置 =================
|
||
UPSTREAM_BODY_LIMIT = 1200 * 1024 # 上游 APIG 请求体限制 ~1.2MB (实测边界1260KB)
|
||
UPSTREAM_TIMEOUT_MIN = 60 # 小请求超时 60s
|
||
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 # 单主机最大连接数
|
||
WAITRESS_THREADS = 64 # Waitress 工作线程数
|
||
SSE_CHUNK_SIZE = 4096 # SSE 流式转发块大小(越小首 token 延迟越低)
|
||
|
||
# ================= 日志 =================
|
||
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_blacklist(self):
|
||
"""清空黑名单"""
|
||
with self._lock:
|
||
self._blacklist.clear()
|
||
|
||
def get_expires_in(self):
|
||
"""返回 token 剩余有效期(秒)"""
|
||
with self._lock:
|
||
return max(0, self._expires_at - time.time())
|
||
|
||
def get_blacklist_count(self):
|
||
"""返回黑名单大小"""
|
||
with self._lock:
|
||
return len(self._blacklist)
|
||
|
||
def fingerprint(self, token):
|
||
"""公开方法:获取 token 指纹"""
|
||
with self._lock:
|
||
return self._fingerprint(token)
|
||
|
||
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"""
|
||
# 快速路径:缓存有效直接返回(避免每次请求都查环境变量/文件)
|
||
cached = cache.get()
|
||
if cached:
|
||
return cached
|
||
|
||
# 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):
|
||
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):
|
||
cache.set(file_token)
|
||
logger.info("Token 从持久化文件加载")
|
||
return file_token
|
||
break
|
||
except (OSError, IOError):
|
||
pass
|
||
|
||
if cache.is_scan_cooldown():
|
||
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 会话池 =================
|
||
# max_retries=0:禁用 urllib3 自动重试,由 proxy 手动控制重试逻辑(避免双重重试)
|
||
http_session = requests.Session()
|
||
adapter = requests.adapters.HTTPAdapter(
|
||
pool_connections=POOL_CONNECTIONS,
|
||
pool_maxsize=POOL_MAXSIZE,
|
||
max_retries=0
|
||
)
|
||
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:
|
||
logger.info(f"上下文估算: ~{total_tokens} tokens (≤{target_limit}), 无需截断 | model={body.get('model','?')} | max_tokens={body.get('max_tokens','?')} | stream={body.get('stream', False)} | messages={len(messages)}")
|
||
# 调试: 超过200KB的请求dump到文件分析
|
||
if len(raw_body) > 200 * 1024:
|
||
import os
|
||
dump_path = f'/tmp/gateway_request_{int(time.time())}.json'
|
||
with open(dump_path, 'wb') as f:
|
||
f.write(raw_body)
|
||
logger.info(f"请求已dump到 {dump_path} ({len(raw_body)//1024}KB)")
|
||
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
|
||
|
||
|
||
# ================= Anthropic API 适配层 (/v1/messages) =================
|
||
|
||
def _anthropic_request_to_openai(body):
|
||
"""Anthropic /v1/messages 请求 → OpenAI /v1/chat/completions 请求
|
||
|
||
转换内容:
|
||
- system (top-level) → messages[0] role=system
|
||
- content blocks: text/image/tool_use/tool_result → OpenAI 格式
|
||
- tools: input_schema → parameters
|
||
- tool_choice: auto/any/tool → auto/required/function
|
||
- stop_sequences → stop
|
||
"""
|
||
openai_body = {}
|
||
openai_body['model'] = body.get('model', '')
|
||
|
||
openai_messages = []
|
||
|
||
# 1. system → system message
|
||
system = body.get('system')
|
||
if system:
|
||
if isinstance(system, str):
|
||
openai_messages.append({'role': 'system', 'content': system})
|
||
elif isinstance(system, list):
|
||
parts = [b.get('text', '') for b in system
|
||
if isinstance(b, dict) and b.get('type') == 'text']
|
||
if parts:
|
||
openai_messages.append({'role': 'system', 'content': '\n'.join(parts)})
|
||
|
||
# 2. 转换 messages
|
||
for msg in body.get('messages', []):
|
||
role = msg.get('role', 'user')
|
||
content = msg.get('content')
|
||
|
||
if isinstance(content, str):
|
||
openai_messages.append({'role': role, 'content': content})
|
||
continue
|
||
|
||
if not isinstance(content, list):
|
||
openai_messages.append({'role': role, 'content': str(content) if content else ''})
|
||
continue
|
||
|
||
# Content blocks 数组
|
||
text_parts = []
|
||
tool_calls = []
|
||
tool_results = []
|
||
has_image = False
|
||
multi_content = []
|
||
|
||
for block in content:
|
||
if not isinstance(block, dict):
|
||
continue
|
||
btype = block.get('type')
|
||
|
||
if btype == 'text':
|
||
text_parts.append(block.get('text', ''))
|
||
multi_content.append({'type': 'text', 'text': block.get('text', '')})
|
||
elif btype == 'image':
|
||
has_image = True
|
||
source = block.get('source', {})
|
||
if source.get('type') == 'base64':
|
||
mt = source.get('media_type', 'image/png')
|
||
multi_content.append({
|
||
'type': 'image_url',
|
||
'image_url': {'url': f'data:{mt};base64,{source.get("data", "")}'}
|
||
})
|
||
elif btype == 'tool_use':
|
||
tool_calls.append({
|
||
'id': block.get('id', ''),
|
||
'type': 'function',
|
||
'function': {
|
||
'name': block.get('name', ''),
|
||
'arguments': json.dumps(block.get('input', {}), ensure_ascii=False)
|
||
}
|
||
})
|
||
elif btype == 'tool_result':
|
||
tool_results.append(block)
|
||
|
||
# tool_result → 独立的 tool 消息
|
||
if tool_results:
|
||
for tr in tool_results:
|
||
tr_content = tr.get('content', '')
|
||
if isinstance(tr_content, list):
|
||
tr_text = '\n'.join(
|
||
b.get('text', '') for b in tr_content
|
||
if isinstance(b, dict) and b.get('type') == 'text'
|
||
)
|
||
else:
|
||
tr_text = str(tr_content) if tr_content else ''
|
||
openai_messages.append({
|
||
'role': 'tool',
|
||
'tool_call_id': tr.get('tool_use_id', ''),
|
||
'content': tr_text
|
||
})
|
||
continue
|
||
|
||
# 构建 assistant/user 消息
|
||
msg_dict = {'role': role}
|
||
if has_image:
|
||
msg_dict['content'] = multi_content
|
||
else:
|
||
msg_dict['content'] = '\n'.join(text_parts) if text_parts else ''
|
||
|
||
if tool_calls:
|
||
msg_dict['tool_calls'] = tool_calls
|
||
if not msg_dict.get('content'):
|
||
msg_dict['content'] = None
|
||
|
||
openai_messages.append(msg_dict)
|
||
|
||
openai_body['messages'] = openai_messages
|
||
|
||
# 3. 参数映射
|
||
openai_body['max_tokens'] = body.get('max_tokens', 4096)
|
||
if 'temperature' in body:
|
||
openai_body['temperature'] = body['temperature']
|
||
if 'top_p' in body:
|
||
openai_body['top_p'] = body['top_p']
|
||
if 'stop_sequences' in body:
|
||
openai_body['stop'] = body['stop_sequences']
|
||
if body.get('stream'):
|
||
openai_body['stream'] = True
|
||
|
||
# metadata.user_id → user
|
||
if body.get('metadata', {}).get('user_id'):
|
||
openai_body['user'] = body['metadata']['user_id']
|
||
|
||
# 4. tools 转换
|
||
if body.get('tools'):
|
||
openai_body['tools'] = [{
|
||
'type': 'function',
|
||
'function': {
|
||
'name': t.get('name', ''),
|
||
'description': t.get('description', ''),
|
||
'parameters': t.get('input_schema', {'type': 'object', 'properties': {}})
|
||
}
|
||
} for t in body['tools']]
|
||
|
||
# 5. tool_choice 转换
|
||
tc = body.get('tool_choice')
|
||
if tc and isinstance(tc, dict):
|
||
tct = tc.get('type', 'auto')
|
||
if tct == 'auto':
|
||
openai_body['tool_choice'] = 'auto'
|
||
elif tct == 'any':
|
||
openai_body['tool_choice'] = 'required'
|
||
elif tct == 'tool':
|
||
openai_body['tool_choice'] = {
|
||
'type': 'function',
|
||
'function': {'name': tc.get('name', '')}
|
||
}
|
||
|
||
return openai_body
|
||
|
||
|
||
def _openai_response_to_anthropic(resp_json, model):
|
||
"""OpenAI 非流式响应 → Anthropic 响应格式
|
||
|
||
转换内容:
|
||
- choices[0].message.content → content[{type:text}]
|
||
- choices[0].message.tool_calls → content[{type:tool_use}]
|
||
- finish_reason → stop_reason (stop→end_turn, length→max_tokens, tool_calls→tool_use)
|
||
- usage.prompt_tokens → usage.input_tokens, usage.completion_tokens → usage.output_tokens
|
||
"""
|
||
choices = resp_json.get('choices', [])
|
||
choice = choices[0] if choices else {}
|
||
message = choice.get('message', {})
|
||
|
||
# 构建 content 数组
|
||
content_blocks = []
|
||
|
||
# 文本内容
|
||
text = message.get('content')
|
||
if text:
|
||
content_blocks.append({'type': 'text', 'text': text})
|
||
|
||
# tool_calls → tool_use blocks
|
||
for tc in message.get('tool_calls', []):
|
||
func = tc.get('function', {})
|
||
try:
|
||
input_data = json.loads(func.get('arguments', '{}'))
|
||
except json.JSONDecodeError:
|
||
input_data = {}
|
||
content_blocks.append({
|
||
'type': 'tool_use',
|
||
'id': tc.get('id', ''),
|
||
'name': func.get('name', ''),
|
||
'input': input_data
|
||
})
|
||
|
||
if not content_blocks:
|
||
content_blocks.append({'type': 'text', 'text': ''})
|
||
|
||
# stop_reason 映射
|
||
fr_map = {
|
||
'stop': 'end_turn',
|
||
'length': 'max_tokens',
|
||
'tool_calls': 'tool_use',
|
||
'content_filter': 'end_turn',
|
||
}
|
||
stop_reason = fr_map.get(choice.get('finish_reason'), 'end_turn')
|
||
|
||
# usage 映射
|
||
usage = resp_json.get('usage', {})
|
||
|
||
return {
|
||
'id': 'msg_' + resp_json.get('id', str(int(time.time() * 1000))),
|
||
'type': 'message',
|
||
'role': 'assistant',
|
||
'model': model,
|
||
'content': content_blocks,
|
||
'stop_reason': stop_reason,
|
||
'stop_sequence': None,
|
||
'usage': {
|
||
'input_tokens': usage.get('prompt_tokens', 0),
|
||
'output_tokens': usage.get('completion_tokens', 0),
|
||
}
|
||
}
|
||
|
||
|
||
def _openai_sse_to_anthropic_sse(resp, model):
|
||
"""OpenAI SSE 流 → Anthropic SSE 流 (生成器,yield bytes)
|
||
|
||
事件序列:
|
||
message_start → content_block_start → content_block_delta* → content_block_stop
|
||
→ (更多 content blocks...) → message_delta → message_stop
|
||
"""
|
||
def sse(event_type, data):
|
||
return f"event: {event_type}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n".encode('utf-8')
|
||
|
||
msg_id = None
|
||
message_started = False
|
||
text_block_open = False
|
||
text_block_index = -1
|
||
next_block_index = 0 # 下一个 content block 的索引
|
||
tool_map = {} # openai_tool_index -> anthropic_block_index
|
||
finish_reason = None
|
||
output_tokens = 0
|
||
input_tokens = 0
|
||
done = False
|
||
buffer = ''
|
||
|
||
for chunk in resp.iter_content(chunk_size=SSE_CHUNK_SIZE):
|
||
if not chunk:
|
||
continue
|
||
buffer += chunk.decode('utf-8', errors='replace')
|
||
|
||
while '\n' in buffer:
|
||
line, buffer = buffer.split('\n', 1)
|
||
line = line.strip()
|
||
if not line or not line.startswith('data:'):
|
||
continue
|
||
|
||
data_str = line[5:].strip()
|
||
|
||
if data_str == '[DONE]':
|
||
done = True
|
||
# 关闭未关闭的 content blocks
|
||
if text_block_open:
|
||
yield sse('content_block_stop', {'type': 'content_block_stop', 'index': text_block_index})
|
||
text_block_open = False
|
||
for idx in sorted(tool_map.values()):
|
||
yield sse('content_block_stop', {'type': 'content_block_stop', 'index': idx})
|
||
tool_map.clear()
|
||
|
||
# message_delta + message_stop
|
||
fr_map = {'stop': 'end_turn', 'length': 'max_tokens',
|
||
'tool_calls': 'tool_use', None: 'end_turn'}
|
||
sr = fr_map.get(finish_reason, 'end_turn')
|
||
yield sse('message_delta', {
|
||
'type': 'message_delta',
|
||
'delta': {'stop_reason': sr, 'stop_sequence': None},
|
||
'usage': {'output_tokens': max(1, output_tokens)}
|
||
})
|
||
yield sse('message_stop', {'type': 'message_stop'})
|
||
continue
|
||
|
||
try:
|
||
data = json.loads(data_str)
|
||
except json.JSONDecodeError:
|
||
continue
|
||
|
||
# usage (有些 provider 在流式 chunk 中包含 usage)
|
||
if 'usage' in data:
|
||
u = data['usage']
|
||
output_tokens = u.get('completion_tokens', output_tokens)
|
||
input_tokens = u.get('prompt_tokens', input_tokens)
|
||
|
||
choices = data.get('choices', [])
|
||
if not choices:
|
||
continue
|
||
choice = choices[0]
|
||
delta = choice.get('delta', {})
|
||
|
||
# 首次 chunk: 发送 message_start
|
||
if not message_started:
|
||
message_started = True
|
||
msg_id = 'msg_' + data.get('id', str(int(time.time() * 1000)))
|
||
yield sse('message_start', {
|
||
'type': 'message_start',
|
||
'message': {
|
||
'id': msg_id, 'type': 'message', 'role': 'assistant',
|
||
'content': [], 'model': model,
|
||
'stop_reason': None, 'stop_sequence': None,
|
||
'usage': {'input_tokens': input_tokens, 'output_tokens': 1}
|
||
}
|
||
})
|
||
|
||
fr = choice.get('finish_reason')
|
||
if fr:
|
||
finish_reason = fr
|
||
|
||
# 文本内容 delta
|
||
cd = delta.get('content')
|
||
if cd is not None and cd != '':
|
||
if not text_block_open:
|
||
text_block_open = True
|
||
text_block_index = next_block_index
|
||
next_block_index += 1
|
||
yield sse('content_block_start', {
|
||
'type': 'content_block_start', 'index': text_block_index,
|
||
'content_block': {'type': 'text', 'text': ''}
|
||
})
|
||
yield sse('content_block_delta', {
|
||
'type': 'content_block_delta', 'index': text_block_index,
|
||
'delta': {'type': 'text_delta', 'text': cd}
|
||
})
|
||
|
||
# tool_calls delta
|
||
tcd = delta.get('tool_calls')
|
||
if tcd:
|
||
for tc in tcd:
|
||
ti = tc.get('index', 0)
|
||
if ti not in tool_map:
|
||
# 关闭 text block(如果开着)
|
||
if text_block_open:
|
||
yield sse('content_block_stop', {'type': 'content_block_stop', 'index': text_block_index})
|
||
text_block_open = False
|
||
tool_map[ti] = next_block_index
|
||
next_block_index += 1
|
||
func = tc.get('function', {})
|
||
yield sse('content_block_start', {
|
||
'type': 'content_block_start', 'index': tool_map[ti],
|
||
'content_block': {
|
||
'type': 'tool_use',
|
||
'id': tc.get('id', f'toolu_{tool_map[ti]}'),
|
||
'name': func.get('name', ''),
|
||
'input': {}
|
||
}
|
||
})
|
||
func = tc.get('function', {})
|
||
args = func.get('arguments', '')
|
||
if args:
|
||
yield sse('content_block_delta', {
|
||
'type': 'content_block_delta', 'index': tool_map[ti],
|
||
'delta': {'type': 'input_json_delta', 'partial_json': args}
|
||
})
|
||
|
||
# 未收到 [DONE] 的兜底关闭
|
||
if not done and message_started:
|
||
if text_block_open:
|
||
yield sse('content_block_stop', {'type': 'content_block_stop', 'index': text_block_index})
|
||
for idx in sorted(tool_map.values()):
|
||
yield sse('content_block_stop', {'type': 'content_block_stop', 'index': idx})
|
||
fr_map = {'stop': 'end_turn', 'length': 'max_tokens',
|
||
'tool_calls': 'tool_use', None: 'end_turn'}
|
||
sr = fr_map.get(finish_reason, 'end_turn')
|
||
yield sse('message_delta', {
|
||
'type': 'message_delta',
|
||
'delta': {'stop_reason': sr, 'stop_sequence': None},
|
||
'usage': {'output_tokens': max(1, output_tokens)}
|
||
})
|
||
yield sse('message_stop', {'type': 'message_stop'})
|
||
|
||
|
||
# ================= Flask 应用 =================
|
||
app = Flask(__name__)
|
||
|
||
|
||
# ================= 全局请求日志(捕获所有请求,包括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": cache.get_expires_in(),
|
||
"blacklisted": cache.get_blacklist_count()
|
||
}, 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)
|
||
# 持久化到文件以便重启后恢复
|
||
env_file = '/etc/huawei-gateway.env'
|
||
try:
|
||
with open(env_file, 'w') as f:
|
||
f.write(f'HUAWEI_TOKEN={token}\n')
|
||
os.chmod(env_file, 0o600)
|
||
except (OSError, IOError):
|
||
pass
|
||
logger.info("Token 已手动注入并持久化")
|
||
return {"status": "ok", "token_fingerprint": cache.fingerprint(token)}, 200
|
||
|
||
|
||
# ================= 通用上游请求 =================
|
||
def _forward_upstream(method, target_url, headers, raw_body, cookies, timeout):
|
||
"""向上游发起请求并返回 response 对象(stream=True)"""
|
||
return http_session.request(
|
||
method=method,
|
||
url=target_url,
|
||
headers=headers,
|
||
data=raw_body,
|
||
cookies=cookies,
|
||
allow_redirects=False,
|
||
timeout=timeout,
|
||
stream=True
|
||
)
|
||
|
||
|
||
def _build_response(resp):
|
||
"""根据上游 response 构建转发给客户端的 Flask Response"""
|
||
# 过滤 hop-by-hop 头和压缩编码头
|
||
skip_headers = {'transfer-encoding', 'content-encoding', 'content-length',
|
||
'connection', 'keep-alive', 'upgrade'}
|
||
response_headers = [(k, v) for k, v in resp.headers.items() if k.lower() not in skip_headers]
|
||
|
||
# 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=SSE_CHUNK_SIZE):
|
||
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()
|
||
|
||
# 记录上游非 200 响应(在读取 content 之后,避免提前消费流)
|
||
if resp.status_code != 200:
|
||
try:
|
||
logger.warning(f"上游返回 {resp.status_code}: {content[:500].decode('utf-8', errors='replace')}")
|
||
except Exception:
|
||
logger.warning(f"上游返回 {resp.status_code}")
|
||
|
||
# 华为云 ModelArts 错误 → OpenAI 标准格式
|
||
if resp.status_code >= 400:
|
||
try:
|
||
err = json.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 = json.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
|
||
)
|
||
|
||
|
||
@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
|
||
headers['Connection'] = 'keep-alive'
|
||
|
||
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"},
|
||
{"id": "glm-5.2", "object": "model", "owned": "zhipu"},
|
||
]
|
||
}
|
||
|
||
# ============ 请求体处理:大小校验 + 动态超时 ============
|
||
raw_body = request.get_data()
|
||
body_size = len(raw_body)
|
||
|
||
# 请求体超过上游限制,直接返回清晰错误
|
||
if body_size > UPSTREAM_BODY_LIMIT:
|
||
logger.warning(f"请求体超限: {body_size // 1024}KB > {UPSTREAM_BODY_LIMIT // 1024}KB")
|
||
return {
|
||
"error": {
|
||
"message": f"请求体过大({body_size // 1024}KB),超过API网关限制({UPSTREAM_BODY_LIMIT // 1024}KB)。请减少请求内容长度。",
|
||
"type": "invalid_request_error",
|
||
"code": "content_too_large",
|
||
"param": None
|
||
}
|
||
}, 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
|
||
elif body_size > 200 * 1024:
|
||
upstream_timeout = 180
|
||
else:
|
||
upstream_timeout = UPSTREAM_TIMEOUT_MIN
|
||
|
||
try:
|
||
resp = _forward_upstream(request.method, target_url, headers, raw_body, request.cookies, upstream_timeout)
|
||
|
||
# 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 = _forward_upstream(request.method, target_url, headers, raw_body, request.cookies, upstream_timeout)
|
||
# 如果新 Token 也 401,清空黑名单避免锁死
|
||
if resp.status_code == 401:
|
||
resp.close()
|
||
logger.warning("新 Token 也 401,清空黑名单避免锁死")
|
||
cache.clear_blacklist()
|
||
new_token2 = find_token_in_memory()
|
||
if new_token2:
|
||
headers['Authorization'] = f'Bearer {new_token2}'
|
||
resp = _forward_upstream(request.method, target_url, headers, raw_body, request.cookies, upstream_timeout)
|
||
|
||
# ============ 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 = _forward_upstream(request.method, target_url, headers, raw_body, request.cookies, upstream_timeout)
|
||
if resp.status_code != 429:
|
||
break
|
||
logger.warning("重试仍返回 429")
|
||
|
||
# ============ 504 超时重试 ============
|
||
if resp.status_code == 504:
|
||
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 = _forward_upstream(request.method, target_url, headers, raw_body, request.cookies, upstream_timeout)
|
||
if resp.status_code != 504:
|
||
break
|
||
|
||
return _build_response(resp)
|
||
|
||
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
|
||
|
||
|
||
@app.route('/v1/messages', methods=['POST'])
|
||
def anthropic_messages():
|
||
"""Anthropic /v1/messages 端点
|
||
|
||
流程: Anthropic 请求 → OpenAI 请求 → 上游转发 → OpenAI 响应 → Anthropic 响应
|
||
复用现有的 token 获取、上下文截断、超时、重试逻辑
|
||
"""
|
||
# 获取 token
|
||
real_token = find_token_in_memory()
|
||
if not real_token:
|
||
logger.error("Anthropic: 未找到 Token")
|
||
return {"type": "error", "error": {"type": "authentication_error", "message": "未找到有效Token"}}, 500
|
||
|
||
# 解析 Anthropic 请求
|
||
anthropic_body = request.get_json(force=True, silent=True)
|
||
if not anthropic_body:
|
||
return {"type": "error", "error": {"type": "invalid_request_error", "message": "无效的请求体"}}, 400
|
||
|
||
model = anthropic_body.get('model', '')
|
||
is_stream = anthropic_body.get('stream', False)
|
||
|
||
# 转换为 OpenAI 格式
|
||
try:
|
||
openai_body = _anthropic_request_to_openai(anthropic_body)
|
||
except Exception as e:
|
||
logger.error(f"Anthropic→OpenAI 请求转换失败: {traceback.format_exc()}")
|
||
return {"type": "error", "error": {"type": "invalid_request_error", "message": f"请求转换失败: {e}"}}, 400
|
||
|
||
openai_body_bytes = json.dumps(openai_body, ensure_ascii=False).encode('utf-8')
|
||
|
||
# 上下文自动截断(复用现有逻辑)
|
||
openai_body_bytes, trim_info = _maybe_trim_context(openai_body_bytes)
|
||
body_size = len(openai_body_bytes)
|
||
|
||
# 请求体大小校验
|
||
if body_size > UPSTREAM_BODY_LIMIT:
|
||
logger.warning(f"Anthropic: 请求体超限 {body_size // 1024}KB")
|
||
return {"type": "error", "error": {"type": "invalid_request_error", "message": f"请求体过大({body_size // 1024}KB)"}}, 413
|
||
|
||
# 构建上游请求头
|
||
headers = {
|
||
'Content-Type': 'application/json',
|
||
'Authorization': f'Bearer {real_token}',
|
||
'Host': TARGET_HOST,
|
||
'Connection': 'keep-alive',
|
||
}
|
||
for k, v in request.headers:
|
||
kl = k.lower()
|
||
if kl not in ('host', 'content-length', 'connection', 'accept-encoding',
|
||
'transfer-encoding', 'authorization', 'content-type',
|
||
'x-api-key', 'anthropic-version', 'anthropic-beta'):
|
||
headers[k] = v
|
||
|
||
target_url = f'https://{TARGET_HOST}/v2/chat/completions'
|
||
|
||
# 动态超时
|
||
if body_size > 800 * 1024:
|
||
upstream_timeout = UPSTREAM_TIMEOUT_MAX
|
||
elif body_size > 200 * 1024:
|
||
upstream_timeout = 180
|
||
else:
|
||
upstream_timeout = UPSTREAM_TIMEOUT_MIN
|
||
|
||
try:
|
||
resp = _forward_upstream('POST', target_url, headers, openai_body_bytes,
|
||
request.cookies, upstream_timeout)
|
||
|
||
# 401 重试(复用 proxy 逻辑)
|
||
if resp.status_code == 401:
|
||
resp.close()
|
||
logger.warning("Anthropic: 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 = _forward_upstream('POST', target_url, headers, openai_body_bytes,
|
||
request.cookies, upstream_timeout)
|
||
if resp.status_code == 401:
|
||
resp.close()
|
||
cache.clear_blacklist()
|
||
new_token2 = find_token_in_memory()
|
||
if new_token2:
|
||
headers['Authorization'] = f'Bearer {new_token2}'
|
||
resp = _forward_upstream('POST', target_url, headers, openai_body_bytes,
|
||
request.cookies, upstream_timeout)
|
||
|
||
# 429 限流重试
|
||
if resp.status_code == 429:
|
||
for i in range(1, RETRY_ON_429 + 1):
|
||
resp.close()
|
||
time.sleep(i * 5)
|
||
logger.warning(f"Anthropic: 429,重试 {i}/{RETRY_ON_429}")
|
||
resp = _forward_upstream('POST', target_url, headers, openai_body_bytes,
|
||
request.cookies, upstream_timeout)
|
||
if resp.status_code != 429:
|
||
break
|
||
|
||
# 504 超时重试
|
||
if resp.status_code == 504:
|
||
for i in range(1, RETRY_ON_504 + 1):
|
||
resp.close()
|
||
time.sleep(i * 3)
|
||
logger.warning(f"Anthropic: 504,重试 {i}/{RETRY_ON_504}")
|
||
resp = _forward_upstream('POST', target_url, headers, openai_body_bytes,
|
||
request.cookies, upstream_timeout)
|
||
if resp.status_code != 504:
|
||
break
|
||
|
||
# 错误处理
|
||
if resp.status_code >= 400:
|
||
content = resp.content
|
||
resp.close()
|
||
try:
|
||
err = json.loads(content)
|
||
msg = err.get('error', {}).get('message', '') or err.get('error_msg', str(err))
|
||
except Exception:
|
||
msg = f"上游返回 {resp.status_code}"
|
||
logger.warning(f"Anthropic: 上游错误 {resp.status_code}: {msg[:200]}")
|
||
return {"type": "error", "error": {"type": "api_error", "message": msg}}, resp.status_code
|
||
|
||
# 流式响应: OpenAI SSE → Anthropic SSE
|
||
if is_stream:
|
||
def stream_gen():
|
||
try:
|
||
for chunk in _openai_sse_to_anthropic_sse(resp, model):
|
||
yield chunk
|
||
finally:
|
||
resp.close()
|
||
|
||
return Response(stream_gen(), status=200, headers={
|
||
'Content-Type': 'text/event-stream',
|
||
'Cache-Control': 'no-cache',
|
||
}, direct_passthrough=True)
|
||
|
||
# 非流式响应: OpenAI JSON → Anthropic JSON
|
||
else:
|
||
content = resp.content
|
||
resp.close()
|
||
try:
|
||
openai_resp = json.loads(content)
|
||
anthropic_resp = _openai_response_to_anthropic(openai_resp, model)
|
||
return Response(
|
||
json.dumps(anthropic_resp, ensure_ascii=False).encode('utf-8'),
|
||
status=200, headers={'Content-Type': 'application/json'}
|
||
)
|
||
except Exception as e:
|
||
logger.error(f"OpenAI→Anthropic 响应转换失败: {traceback.format_exc()}")
|
||
return {"type": "error", "error": {"type": "api_error", "message": f"响应转换失败: {e}"}}, 500
|
||
|
||
except requests.exceptions.Timeout:
|
||
logger.error("Anthropic: 请求超时")
|
||
return {"type": "error", "error": {"type": "api_error", "message": "请求超时"}}, 504
|
||
except requests.exceptions.ConnectionError:
|
||
logger.error("Anthropic: 无法连接上游")
|
||
return {"type": "error", "error": {"type": "api_error", "message": "无法连接上游服务"}}, 502
|
||
except Exception as e:
|
||
logger.error(f"Anthropic 端点错误: {traceback.format_exc()}")
|
||
return {"type": "error", "error": {"type": "api_error", "message": 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=WAITRESS_THREADS)
|
||
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()
|