148 lines
6.2 KiB
Python
148 lines
6.2 KiB
Python
"""pipeline_llm.selection — 模型选择的唯一收敛点(下拉数据/引用解析/缺省模型)。
|
||
|
||
2026-09-04 收敛改造:此前模型选择逻辑散在 pipeline-sdlc(cockpit/get_model_options/
|
||
set_agent_model)、pipeline-core(agent_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")
|
||
|
||
|
||
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'):
|
||
"""模型选择下拉的唯一数据源(对齐 llmage 分类:capability=能力类型)。
|
||
|
||
机构语义与推理链一致:本机构 + 系统级共享(org_id 空/'0')可见。
|
||
selected 标记:项目已设模型(sd_projects.default_model,存 name)>
|
||
个人全局选择(pipeline_agent_settings.default_llm_id,存 id)。
|
||
value_field: 'id'=下拉值用模型 id(AgentIO 场景);'name'=用注册名(角色模型配置场景)。
|
||
"""
|
||
db, dbname = _get_sor()
|
||
rows = []
|
||
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
|
||
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=''):
|
||
"""模型引用解析的唯一入口:id / vendor_model_id / name → 模型注册名。
|
||
|
||
机构隔离:非系统级机构只能解析「本机构 + 系统级共享」模型。
|
||
解析不到返回 ''(调用方决定回退/报错,禁止自行另查表)。
|
||
"""
|
||
if not model_ref:
|
||
return ''
|
||
db, dbname = _get_sor()
|
||
async with db.sqlorContext(dbname) as sor:
|
||
recs = await sor.sqlExe(
|
||
"SELECT name, org_id 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 ''
|
||
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")
|