refactor(inference): 组包借鉴uapi模板渲染——path/headers/body全Jinja2(tmpl_engine.renders+json.loads),api_key走模板命名空间注入;废弃土制字符串替换(401根因方向);内置openai-chat缺省模板,新形态只加profile不改代码

This commit is contained in:
yumoqing 2026-09-04 18:21:10 +08:00
parent e0c1dc751d
commit 26d083bdd2

View File

@ -19,9 +19,20 @@ import logging
logger = logging.getLogger("pipeline_llm.inference")
# OpenAI 兼容缺省适配(模型未挂 profile / profile 缺字段时的兜底)
# OpenAI 兼容缺省适配(模型未挂 profile / profile 缺字段时的兜底)——
# 组包方式借鉴 uapi 模块llmage 即调 uapi 组包path/headers/body 全部
# Jinja2 模板,经 tmpl_engine 渲染后 json.loadsapi_key 作模板变量注入。
_DEFAULT_CHAT_PATH = "/chat/completions"
_DEFAULT_CHAT_HEADERS = '{"Authorization": "***}"}'
_BUILTIN_HEADERS_TMPL = '{"Authorization": "***", "Content-Type": "application/json"}'
# 缺省请求体模板model+messages 必填,其余调用方参数经 extra 透传
_BUILTIN_REQUEST_TMPL = (
'{"model": {{ model|tojson }}, "messages": {{ messages|tojson }}'
'{% for k, v in extra.items() %}, "{{ k }}": {{ v|tojson }}{% endfor %}}'
)
# 请求体透传字段白名单OpenAI 兼容常用参数,除 model/messages 外)
_PASSTHROUGH_FIELDS = ('temperature', 'top_p', 'max_tokens', 'tools',
'tool_choice', 'presence_penalty', 'frequency_penalty',
'stop', 'response_format', 'seed')
_RETRYABLE_STATUS = (429, 500, 502, 503, 504)
_MAX_ATTEMPTS = 3
_TOTAL_TIMEOUT = 300
@ -98,7 +109,7 @@ async def _resolve_call(sor, org_id, user_id, model_name, capability,
raise GovernError(
'模型「%s」能力为 %s,与请求能力 %s 不匹配。请按能力分类选择模型' % (
model_row.get('name', ''), cap, capability))
# 适配模板(决定上游 path/headers
# 适配模板(组包path/headers/request_template全 Jinja2借鉴 uapi
profile = {}
pid = model_row.get('profile_id') or ''
if pid:
@ -111,6 +122,8 @@ async def _resolve_call(sor, org_id, user_id, model_name, capability,
profile = {
'path': getattr(r, 'path', '') or '',
'headers': getattr(r, 'headers', '') or '',
'request_template': getattr(r, 'request_template', '') or '',
'response_template': getattr(r, 'response_template', '') or '',
}
# default_params 合并(调用方显式参数优先)
ctx['default_params'] = _parse_json(model_row.get('default_params'), {})
@ -156,52 +169,122 @@ async def _pick_default_model_name(org_id, capability='t2t'):
return ''
# ────────────────────────── 上游调用 ──────────────────────────
# ────────────────────────── 上游调用(组包借鉴 uapi ──────────────────────────
# uapi 模块llmage 即调它组包的做法path/headers/body 全是 Jinja2 模板,
# tmpl_engine.renders() 渲染后 json.loadsapi_key 等作模板变量注入命名空间。
# 本层同款实现:模板来自 llm_api_profile模型挂的适配模板缺省用内置
# openai-chat 模板——新增供应商形态只加 profile 记录,不改代码。
def _render_headers(tmpl, api_key):
"""适配模板 headersJSON 对象,值里 {api_key} 占位符替换为真实 key。"""
h = _parse_json(tmpl, None)
if not isinstance(h, dict) or not h:
h = _parse_json(_DEFAULT_CHAT_HEADERS, {"Authorization": "***}"})
out = {}
for k, v in h.items():
v = str(v)
if '{api_key}' in v:
v = v.replace('{api_key}', api_key)
out[str(k)] = v
# 兜底认证头(模板没写 Authorization 时补标准 Bearer
_tmpl_engine_cache = None
async def _render_tmpl(tmplstr, ns):
"""渲染一段模板字符串uapi rendertmpl 同款)。
服务进程用 ahserver tmpl_engine启动时 setupTemplateEngine 注入异步
独立脚本环境没有引擎时用 jinja2.Environment 兜底同样异步渲染
"""
global _tmpl_engine_cache
te = _tmpl_engine_cache
if te is None:
try:
from ahserver.serverenv import ServerEnv
te = getattr(ServerEnv(), 'tmpl_engine', None)
except Exception:
te = None
if te is None or not callable(getattr(te, 'renders', None)):
try:
from jinja2 import Environment
te = Environment(enable_async=True)
except Exception:
raise GovernError('模板引擎不可用jinja2 未安装或未初始化,无法组包上游请求')
_tmpl_engine_cache = te
return await te.renders(tmplstr, ns)
async def _render_headers(tmpl, ns):
"""headers 模板 → dictJinja2 渲染后 json.loads同 uapi
模板空用内置 Bearer 模板渲染结果缺 Authorization 头时兜底补标准 Bearer
"""
s = await _render_tmpl((tmpl or '').strip() or _BUILTIN_HEADERS_TMPL, ns)
try:
out = json.loads(s)
except Exception as e:
raise GovernError('请求头模板渲染结果不是合法 JSON%s(模板:%s' % (
str(e)[:120], (tmpl or '(内置)')[:120]))
if not isinstance(out, dict):
raise GovernError('请求头模板必须渲染成 JSON 对象')
if not any(str(k).lower() == 'authorization' for k in out):
out['Authorization'] = 'Bearer %s' % api_key
out['Authorization'] = 'Bearer %s' % (ns.get('api_key') or '')
out.setdefault('Content-Type', 'application/json')
return out
async def _call_upstream_chat(ctx, payload):
"""POST 上游 chat/completions按适配模板 path/headers带瞬时错误重试。
async def _build_upstream_body(ctx, payload):
"""请求体模板渲染uapi data 模板同款)
返回上游响应 dict失败抛 GovernError/ValueError消息真实可行动
合并顺序模型 default_params 打底 调用方显式参数覆盖 模板渲染
模板取 profile.request_template缺省内置 openai-chat 模板
模板命名空间model / messages / extra(白名单透传参数) / params(全部) /
api_key / org_id
"""
profile = ctx.get('profile') or {}
tmpl = (profile.get('request_template') or '').strip() or _BUILTIN_REQUEST_TMPL
upstream = dict(payload or {})
upstream.pop('stream', None) # 一期不透传流式
upstream['model'] = ctx.get('model_id') or upstream.get('model') or ''
for k, v in (ctx.get('default_params') or {}).items():
upstream.setdefault(k, v)
messages = upstream.pop('messages', [])
model = upstream.pop('model', '')
extra = {k: upstream[k] for k in upstream if k in _PASSTHROUGH_FIELDS}
ns = {
'model': model, 'messages': messages, 'extra': extra,
'params': upstream,
'api_key': ctx.get('api_key') or '',
'org_id': ctx.get('org_id') or '',
}
s = await _render_tmpl(tmpl, ns)
try:
body = json.loads(s)
except Exception as e:
raise GovernError('请求体模板渲染结果不是合法 JSON%s(模板:%s' % (
str(e)[:120], tmpl[:120]))
if not isinstance(body, dict):
raise GovernError('请求体模板必须渲染成 JSON 对象')
return body
async def _call_upstream_chat(ctx, payload):
"""POST 上游path/headers/body 全模板组包),带瞬时错误重试。
返回上游响应 dict失败抛 ValueError消息真实可行动
"""
import aiohttp
import asyncio
url = (ctx.get('api_base') or '').rstrip('/') + (
(ctx.get('profile') or {}).get('path') or _DEFAULT_CHAT_PATH)
headers = _render_headers((ctx.get('profile') or {}).get('headers', ''),
ctx.get('api_key') or '')
profile = ctx.get('profile') or {}
ns = {'api_key': ctx.get('api_key') or '', 'org_id': ctx.get('org_id') or ''}
path = await _render_tmpl((profile.get('path') or '').strip() or _DEFAULT_CHAT_PATH, ns)
if path.startswith('http://') or path.startswith('https://'):
url = path
else:
url = (ctx.get('api_base') or '').rstrip('/') + path
headers = await _render_headers(profile.get('headers', ''), ns)
body = await _build_upstream_body(ctx, payload)
timeout = int(_fnum(ctx.get('timeout')) or _TOTAL_TIMEOUT)
upstream = dict(payload or {})
upstream.pop('stream', None) # 一期代理不支持流式透传
upstream['model'] = ctx.get('model_id') or upstream.get('model') or ''
for k, v in (ctx.get('default_params') or {}).items():
upstream.setdefault(k, v)
# 可诊断性:只记长度不记值(真 key 不落日志)
logger.info("inference: 组包 url=%s auth头长度=%d body字段=%s",
url.split('?')[0], len(headers.get('Authorization', '') or ''),
sorted(body.keys()))
last = None
for attempt in range(_MAX_ATTEMPTS):
try:
async with aiohttp.ClientSession() as session:
async with session.post(
url, headers=headers, json=upstream,
url, headers=headers, json=body,
timeout=aiohttp.ClientTimeout(total=timeout, connect=_CONNECT_TIMEOUT),
) as resp:
if resp.status != 200: