xtc-backend / app /api /_common.py
a3216's picture
sync from GitHub 6e39eaf: feat: 添加 latin1_safe_header 函数,确保 HTTP header 值的安全性,避免编码错误
04b9a37 verified
Raw
History Blame Contribute Delete
9.31 kB
"""共享请求/响应工具: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