#!/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/', methods=['POST', 'GET', 'OPTIONS', 'PUT', 'DELETE']) @app.route('/v2/', 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()