"""共享请求/响应工具:CORS、provider 选择辅助、usage 记录。""" from __future__ import annotations import json import logging import time import uuid from typing import Any, Optional from fastapi import Request from fastapi.responses import JSONResponse from ..config import get_settings from ..cors import build_cors_headers from ..database import get_conn from ..models.config import AppConfig, Provider from ..providers.keypool import choose_provider_api_key from ..providers.policy import assert_model_allowed from ..providers.resolver import normalize_model, resolve_provider logger = logging.getLogger(__name__) # 通配模式下的默认 CORS 头(白名单模式下为空,由 ok_json(origin=) 动态计算) # 注:此常量在模块加载时计算一次。若运行时切换 XTC_CORS_ORIGINS,需要重启进程。 CORS_HEADERS = build_cors_headers(None) if CORS_HEADERS and "Access-Control-Expose-Headers" not in CORS_HEADERS: # 通配模式下 build_cors_headers 不带 Expose-Headers,这里补上以保持向后兼容 CORS_HEADERS["Access-Control-Expose-Headers"] = ( "x-xtc-provider, x-xtc-model, x-xtc-image-fix-mode" ) def ok_json( payload: dict, *, status: int = 200, extra_headers: Optional[dict] = None, origin: Optional[str] = None, ) -> JSONResponse: """构造 JSON 响应并附加 CORS 头。 Args: origin: 请求头中的 Origin 值。传入时按白名单动态匹配; 不传时使用模块级 ``CORS_HEADERS``(通配模式才回 ``*``)。 """ if origin is not None: headers = build_cors_headers(origin) if headers and "Access-Control-Expose-Headers" not in headers: headers["Access-Control-Expose-Headers"] = ( "x-xtc-provider, x-xtc-model, x-xtc-image-fix-mode" ) else: headers = dict(CORS_HEADERS) if extra_headers: headers.update(extra_headers) return JSONResponse(status_code=status, content=payload, headers=headers) def ok_with_cors( payload: dict, *, status: int = 200, extra_headers: Optional[dict] = None, origin: Optional[str] = None, ) -> JSONResponse: body = {"ok": True, **payload} return ok_json(body, status=status, extra_headers=extra_headers, origin=origin) def latin1_safe_header(value: Optional[str]) -> str: """把字符串转为 latin-1 安全的 HTTP header 值。 HTTP header 值只支持 latin-1 编码,含非 latin-1 字符(如中文)会导致 "'latin-1' codec can't encode characters in position ..." 500 错误。 对含非 latin-1 字符的值做百分号编码,纯 ASCII 值原样返回。 """ if not value: return "" s = str(value) try: s.encode("latin-1") return s except UnicodeEncodeError: from urllib.parse import quote return quote(s) def gen_trace_id() -> str: return uuid.uuid4().hex[:12] def audit_admin(action: str, admin_key: str, *, target: Optional[str] = None, **detail) -> None: """记录 admin 跨用户数据操作审计日志(NC3)。 actor 用 admin key 的 sha256 前 12 位(避免明文 key 入库但仍可区分不同 admin)。 """ import hashlib from ..services import usage_store actor = "admin" if admin_key: actor = "admin:" + hashlib.sha256(admin_key.encode("utf-8")).hexdigest()[:12] usage_store.audit(action=action, actor=actor, target=target, detail=detail or None) async def select_provider_and_key( *, config: AppConfig, provider_id: Optional[str], model: Optional[str], fallback_model: Optional[str] = None, ) -> tuple[Provider, str, str]: """统一流程:解析 provider -> 校验策略 -> 选 key -> 归一化模型名。""" provider = resolve_provider(config, provider_id=provider_id, model=model) clean_model = normalize_model(provider, model, fallback=fallback_model) assert_model_allowed(provider, clean_model) api_key = choose_provider_api_key(provider) return provider, api_key, clean_model def record_usage( *, access_key: Optional[str], provider: str, model: str, usage: Optional[dict], ok: bool, error_code: Optional[str] = None, ) -> None: now = int(time.time()) sql = ( "INSERT INTO usage_log(ts, access_key, provider, model, prompt_tokens, completion_tokens, total_tokens, ok, error_code) " "VALUES(?,?,?,?,?,?,?,?,?)" ) args = ( now, access_key, provider, model, int((usage or {}).get("prompt_tokens") or 0), int((usage or {}).get("completion_tokens") or 0), int((usage or {}).get("total_tokens") or 0), 1 if ok else 0, error_code, ) def _task(c): try: c.execute(sql, args) except Exception as e: logger.warning("[usage] record failed: %s", e) # 转交后台写线程,避免在事件循环上同步 INSERT try: from ..db_writer import enqueue, is_started if is_started(): enqueue(_task) else: with get_conn() as conn: _task(conn) except Exception: try: with get_conn() as conn: _task(conn) except Exception as e: logger.warning("[usage] record failed: %s", e) # Webhook 通知(fire-and-forget,失败不影响主流程) try: from ..services import webhook_store event = "chat.completed" if ok else "chat.failed" webhook_store.notify_fire_and_forget( event, { "access_key": access_key, "provider": provider, "model": model, "ok": ok, "error_code": error_code, "usage": usage, "ts": int(time.time()), }, ) except Exception: pass async def read_json_body(request: Request) -> dict: """读取 JSON body,失败返回空 dict(兼容 multipart 场景)。""" try: data = await request.json() if isinstance(data, dict): return data except Exception: pass return {} def normalize_chat_body(body: dict) -> tuple[list, str, str]: """把 XTC 简化 body 归一化为 (messages, model, provider_id)。 兼容前端 api.js 的两种调用形态: - 简化形态(默认):``{input: "文本", images: ["url"...], files: [...]}`` 自动组装为 OpenAI messages,等价于旧 Netlify 版 xtc-client-api.mjs 的行为。 - OpenAI 形态:``{messages: [{role, content}]}`` 直接透传。 同时处理附件文本(files/file_names),追加为最后一条 user message。 file_names 兼容 JSON 字符串(前端 upload 路径会发 JSON.stringify(names))。 """ if not isinstance(body, dict): return [], "", "" messages = body.get("messages") if not isinstance(messages, list): messages = [] input_text = str(body.get("input") or "").strip() images = body.get("images") if not isinstance(images, list): images = [images] if images else [] image_urls: list[str] = [] for u in images: s = str(u or "").strip() if s: image_urls.append(s) # 简化形态:只有 input/images 时,组装成 OpenAI messages if not messages and (input_text or image_urls): content: list = [] if input_text: content.append({"type": "text", "text": input_text}) for url in image_urls: content.append({"type": "image_url", "image_url": {"url": url}}) messages = [{"role": "user", "content": content}] if not isinstance(messages, list): messages = [] # 附件文本拼到末尾(前端 normalizeFileInputs 上传的文本类文件) files = body.get("files") if not isinstance(files, list): files = [files] if files else [] file_names_raw = body.get("file_names") # file_names 兼容 JSON 字符串(前端 callUpload 会发 JSON.stringify(names)) if isinstance(file_names_raw, str) and file_names_raw.strip().startswith("["): try: parsed = json.loads(file_names_raw) if isinstance(parsed, list): file_names_raw = [str(x) for x in parsed] except Exception: file_names_raw = [file_names_raw] if not isinstance(file_names_raw, list): file_names_raw = [file_names_raw] if file_names_raw else [] text_parts: list[str] = [] for i, f in enumerate(files): if not isinstance(f, str): f = str(f or "") if not f.strip(): continue fname = "" if i < len(file_names_raw): fname = str(file_names_raw[i] or "").strip() fname = fname or f"file_{i + 1}" text_parts.append(f"[{fname}]\n{f}") if text_parts: messages = list(messages) + [ {"role": "user", "content": "\n\n".join(text_parts)} ] model = str(body.get("model") or "").strip() provider_id = body.get("provider") or body.get("provider_id") or "" provider_id = str(provider_id or "").strip() return messages, model, provider_id