pipeline-llm/pipeline_llm/selection.py

148 lines
6.2 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")
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'=下拉值用模型 idAgentIO 场景);'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")