pipeline-llm/pipeline_llm/selection.py
yumoqing 8e32cb85f3 feat(llm): 会话agent可选模型收窄为对话能力t2t/i2t/m2t+能力字典加m2t(2026-09-06用户定夺)
- selection.CHAT_CAPS=('t2t','i2t','m2t') 作为会话形态能力白名单唯一事实源
- model_options/resolve_model_name 加 capabilities 参数('chat'哨兵),
  IN 列表展开占位符(sqlor传list会崩),capability 空串按 t2t 归一(COALESCE+NULLIF)
- chat_inference 同步门禁从「只认 t2t」放宽到 CHAT_CAPS——
  实测根因:测试库 6 个机构策略主模型全是 qwen3.8-max(i2t),
  16:28 已产生 FAILED「capability mismatch: i2t != t2t」,会话agent选它必挂
- m2t 入字典四处同步:种子 init/data.json + 端点注释 models.dspy
  + design-spec §6 能力表(+ 提取提示词在 pipeline-platform)
2026-09-06 19:32:59 +08:00

189 lines
8.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""pipeline_llm.selection — 模型选择的唯一收敛点(下拉数据/引用解析/缺省模型)。
2026-09-04 收敛改造:此前模型选择逻辑散在 pipeline-sdlccockpit/get_model_options/
set_agent_model、pipeline-coreagent_model_options、旧 llm CRUD、pipeline-service
gateway/agent_loop 内联 SQL多处各查各的表、语义漂移。现全部收敛到本模块
- model_options() 所有模型下拉的唯一数据源
- resolve_model_name() 模型引用id/vendor_model_id/name→ 注册名的唯一解析点
- org_available_models() 机构可用模型清单(冒泡检测用)
- 缺省模型由机构策略决定inference._pick_default_model_name不再写死模型名
其他模块只允许薄壳委托dspy 3-5 行)或函数调用,禁止再内联模型表 SQL。
"""
import logging
from sqlor.dbpools import DBPools
from ahserver.serverenv import ServerEnv
logger = logging.getLogger("pipeline_llm.selection")
# 会话/chat 形态可用的能力类型2026-09-06 用户定夺):文本输出的对话能力。
# 会话 agent 模型下拉只列这些embedding/rerank/图视频生成等形态各有专属入口。
CHAT_CAPS = ('t2t', 'i2t', 'm2t')
def _norm_caps(capabilities):
"""capabilities 参数归一:'chat' 哨兵 → CHAT_CAPS逗号串 → tuple
tuple/list 原样。返回 tuple空 = 不过滤。dspy 薄壳只传 'chat',零 import。"""
if isinstance(capabilities, str):
s = capabilities.strip().lower()
if not s:
return ()
if s == 'chat':
return CHAT_CAPS
return tuple(c.strip().lower() for c in s.split(',') if c.strip())
return tuple(str(c).strip().lower() for c in (capabilities or ()) if str(c or '').strip())
def _get_sor():
env = ServerEnv()
fn = getattr(env, 'get_module_dbname', None)
dbname = 'pipeline'
if callable(fn):
try:
dbname = fn('pipeline_llm') or 'pipeline'
except Exception:
dbname = 'pipeline'
return DBPools(), dbname
async def model_options(org_id, uid='', session_id='', pipeline_id='',
value_field='id', capabilities=()):
"""模型选择下拉的唯一数据源(对齐 llmage 分类capability=能力类型)。
机构语义与推理链一致:本机构 + 系统级共享org_id 空/'0')可见。
selected 标记项目已设模型sd_projects.default_model存 name>
个人全局选择pipeline_agent_settings.default_llm_id存 id
value_field: 'id'=下拉值用模型 idAgentIO 场景);'name'=用注册名(角色模型配置场景)。
capabilities: 能力类型白名单tuple/list如会话 agent 传 CHAT_CAPS=('t2t','i2t','m2t'))。
空 = 不过滤(历史行为)。注意:存量模型 capability 为空的行按 't2t' 语义对待,
过滤时用 COALESCE+NULLIF 把空串归一成 't2t' 再比对DDL 有默认值,
但历史迁移行可能是空串)。
"""
db, dbname = _get_sor()
rows = []
caps = _norm_caps(capabilities)
async with db.sqlorContext(dbname) as sor:
sql = ("SELECT m.id, m.name, m.vendor_model_id, m.capability, v.name AS vendor_name "
"FROM llm_model m LEFT JOIN llm_vendor v ON v.id=m.vendor_id "
"WHERE m.status='active'")
params = {}
if org_id and org_id != '0':
sql += " AND (m.org_id=${org}$ OR m.org_id='' OR m.org_id='0')"
params['org'] = org_id
if caps:
# sqlor 的 IN 列表必须展开占位符(传 list 会崩):${c0}$,${c1}$,...
ph = []
for i, c in enumerate(caps):
k = 'c%d' % i
ph.append('${%s}$' % k)
params[k] = c
sql += (" AND COALESCE(NULLIF(m.capability,''),'t2t') IN (%s)"
% ','.join(ph))
sql += " ORDER BY m.name"
recs = await sor.sqlExe(sql, params)
# 项目已设模型(按会话解析当前项目)
project_model = ''
if uid:
try:
from pipeline_service.workspace import get_session_project_id
_pid = await get_session_project_id(
sor, uid, session_id or '', pipeline_id or '')
if _pid:
_p = await sor.sqlExe(
"SELECT default_model FROM sd_projects WHERE id=${p}$", {"p": _pid})
await sor.sqlExe("COMMIT", {})
if _p:
project_model = getattr(_p[0], 'default_model', '') or ''
except Exception as e:
logger.debug("model_options 项目模型解析跳过: %s", e)
# 个人全局默认选择
current_llm_id = ''
if uid:
try:
_s = await sor.sqlExe(
"SELECT default_llm_id FROM pipeline_agent_settings WHERE user_id=${u}$",
{"u": uid})
if _s:
current_llm_id = getattr(_s[0], 'default_llm_id', '') or ''
except Exception:
pass
for r in (recs or []):
vname = getattr(r, 'vendor_name', '') or ''
rows.append({
'value': r.id if value_field != 'name' else r.name,
'text': r.name + (' (' + vname + ')' if vname else ''),
'provider': vname,
'model_id': r.id,
'model_id_text': r.name,
'capabilities': getattr(r, 'capability', '') or 't2t',
'selected': ((r.name == project_model) if project_model
else (r.id == current_llm_id)),
})
return rows
async def resolve_model_name(model_ref, org_id='', capabilities=()):
"""模型引用解析的唯一入口id / vendor_model_id / name → 模型注册名。
机构隔离:非系统级机构只能解析「本机构 + 系统级共享」模型。
capabilities: 能力白名单('chat' 哨兵 / tuple非空时能力不符也解析失败
(会话 agent 个人默认模型/角色模型入口传 'chat',防把 embedding 等非对话
模型持久化成会话模型)。
解析不到返回 ''(调用方决定回退/报错,禁止自行另查表)。
"""
if not model_ref:
return ''
caps = _norm_caps(capabilities)
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
recs = await sor.sqlExe(
"SELECT name, org_id, capability FROM llm_model WHERE status='active' "
"AND (id=${m}$ OR vendor_model_id=${m}$ OR name=${m}$) LIMIT 1",
{"m": model_ref})
await sor.sqlExe("COMMIT", {})
if not recs:
return ''
m_org = getattr(recs[0], 'org_id', '') or ''
if org_id and org_id != '0' and m_org not in ('', '0', org_id):
logger.warning("resolve_model_name: 模型 %s 不属于机构 %s", model_ref, org_id)
return ''
if caps:
cap = (getattr(recs[0], 'capability', '') or 't2t').strip().lower()
if cap not in caps:
logger.warning("resolve_model_name: 模型 %s 能力 %s 不在白名单 %s",
model_ref, cap, caps)
return ''
return getattr(recs[0], 'name', '') or ''
async def org_available_models(org_id, capability=''):
"""机构可用模型注册名列表(本机构 + 系统级共享)。空 = 机构未配置模型。"""
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
sql = "SELECT name, capability FROM llm_model WHERE status='active'"
params = {}
if org_id and org_id != '0':
sql += " AND (org_id=${org}$ OR org_id='' OR org_id='0')"
params['org'] = org_id
if capability:
sql += " AND capability=${c}$"
params['c'] = capability
recs = await sor.sqlExe(sql, params)
await sor.sqlExe("COMMIT", {})
return [getattr(r, 'name', '') or '' for r in (recs or [])]
def load_selection():
"""注册到 ServerEnv供各模块 dspy 薄壳调用)。"""
env = ServerEnv()
env.llm_model_options = model_options
env.llm_resolve_model_name = resolve_model_name
env.llm_org_available_models = org_available_models
logger.info("[pipeline_llm] selection loaded")