508 lines
22 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 模块入口 — 模型治理(供应商/账号/模型/组织策略/限额限流/记账)。
宿主集成:应用入口调用 load_llm()try/except ImportError 兜底)。
所有表共享宿主应用的库get_module_dbname('pipeline_llm')),表名统一 llm_ 前缀。
模块只提供治理引擎与配置管理;实际 LLM 转发仍由宿主已有的
pipeline_service.llm_bridge / llm_proxy 承担——本模块在两个调用链上挂钩子:
- llm_bridge._get_model_config治理启用时前置解析主备容错+端点选择)
- llm_proxy.proxy_chat_completion调用前后接门禁与结算运行环境路径
"""
import json
import logging
from sqlor.dbpools import DBPools
from ahserver.serverenv import ServerEnv
from appPublic.uniqueID import getID
from .gateway import (
govern_resolve, govern_settle, governance_enabled,
recharge_account, recharge_org,
encrypt_api_key, decrypt_api_key, load_gateway, GovernError,
)
logger = logging.getLogger("pipeline_llm")
MODULE_NAME = 'pipeline_llm'
def _dbname():
"""动态取宿主库名(禁硬编码)。未注册 get_module_dbname 时回落 pipeline。"""
env = ServerEnv()
fn = getattr(env, 'get_module_dbname', None)
if callable(fn):
try:
return fn(MODULE_NAME) or 'pipeline'
except Exception:
pass
return 'pipeline'
def _get_sor():
return DBPools(), _dbname()
# ────────────────────── 端点目录输入归一化 ──────────────────────
def normalize_endpoints(text):
"""端点目录输入归一化:支持「一行一个 URL」或 JSON 数组两种写法。
行格式URL [region=domestic|international] [timeout=秒],例如:
https://dashscope.aliyuncs.com/compatible-mode/v1
https://api.example.com/v1 region=international timeout=30
缺省 region=domestic、timeout=60。存储格式恒为 JSON 数组
[{base_url, region, timeout}](治理链按此读取)。
非法输入抛 ValueError中文提示直接展示给用户
"""
s = (text or '').strip()
if not s:
return ''
if s.startswith('['):
eps = json.loads(s)
if not isinstance(eps, list):
raise ValueError('endpoints 必须是 JSON 数组')
else:
eps = []
for ln_no, ln in enumerate([l.strip() for l in s.splitlines()], 1):
if not ln:
continue
parts = ln.split()
url = parts[0].rstrip('/')
if not (url.startswith('http://') or url.startswith('https://')):
raise ValueError('%d 行不是合法 URL须以 http(s):// 开头):%s' % (ln_no, parts[0]))
ep = {'base_url': url, 'region': 'domestic', 'timeout': 60}
for kv in parts[1:]:
if '=' not in kv:
raise ValueError('%d 行选项 %s 应为 key=value 形式region= / timeout=' % (ln_no, kv))
k, v = kv.split('=', 1)
if k == 'region':
if v not in ('domestic', 'international'):
raise ValueError('%d 行 region 只能是 domestic 或 international' % ln_no)
ep['region'] = v
elif k == 'timeout':
ep['timeout'] = int(v)
else:
raise ValueError('%d 行未知选项 %s(支持 region= / timeout=' % (ln_no, k))
eps.append(ep)
for i, ep in enumerate(eps):
if not isinstance(ep, dict) or not (ep.get('base_url') or '').strip():
raise ValueError('端点 #%d 缺少 base_url' % (i + 1))
return json.dumps(eps, ensure_ascii=False)
def _norm_endpoint_ids(raw):
"""endpoint_ids 归一化JSON 数组或逗号分隔数字均可,返回 JSON 数组串。"""
if isinstance(raw, (list, dict)):
raw = json.dumps(raw, ensure_ascii=False)
raw = (raw or '').strip()
if not raw:
return ''
if raw.startswith('['):
idxs = json.loads(raw)
else:
idxs = [x.strip() for x in raw.split(',') if x.strip()]
if not isinstance(idxs, list):
raise ValueError('endpoint_ids 必须是 JSON 数组')
for i in idxs:
if not (isinstance(i, int) or (isinstance(i, str) and i.lstrip('-').isdigit())):
raise ValueError('endpoint_ids 元素必须是端点下标(数字)')
return json.dumps([int(i) for i in idxs], ensure_ascii=False)
MASK = '******' # 列表脱敏占位;编辑表单原样提交视为「不修改」
# ────────────────────── CRUD供生成的 CRUD 页面调用) ──────────────────────
def _clean(params_kw):
data = dict(params_kw or {})
for k in ('page', 'rows', 'data_filter', 'sortby'):
data.pop(k, None)
return data
async def create_llm_vendor(params_kw):
result = {'success': False, 'message': ''}
try:
data = _clean(params_kw)
data['endpoints'] = normalize_endpoints(data.get('endpoints', ''))
data['id'] = getID()
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
await sor.C('llm_vendor', data)
result['success'] = True
result['message'] = '创建成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def update_llm_vendor(params_kw):
result = {'success': False, 'message': ''}
try:
data = _clean(params_kw)
if 'endpoints' in data:
data['endpoints'] = normalize_endpoints(data.get('endpoints', ''))
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
await sor.U('llm_vendor', data)
result['success'] = True
result['message'] = '更新成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def delete_llm_vendor(params_kw):
"""供应商删除前置检查:有账号/模型引用时禁止(防孤儿)。"""
result = {'success': False, 'message': ''}
try:
vid = (params_kw or {}).get('id', '')
if not vid:
raise ValueError('缺少 id')
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
recs = await sor.sqlExe(
"SELECT (SELECT COUNT(*) FROM llm_account WHERE vendor_id=${v}$) + "
"(SELECT COUNT(*) FROM llm_model WHERE vendor_id=${v}$) AS c", {"v": vid})
await sor.sqlExe("COMMIT", {})
cnt = int(getattr(recs[0], 'c', 0)) if recs else 0
if cnt > 0:
raise ValueError('该供应商下还有 %d 个账号/模型,请先移除后再删除' % cnt)
await sor.sqlExe("DELETE FROM llm_vendor WHERE id=${v}$", {"v": vid})
await sor.sqlExe("COMMIT", {})
result['success'] = True
result['message'] = '删除成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def create_llm_account(params_kw):
"""账号创建api_key 自动 AES 加密;校验端点下标合法。"""
result = {'success': False, 'message': ''}
try:
data = _clean(params_kw)
vid = data.get('vendor_id', '')
if not vid:
raise ValueError('必须选择供应商')
# api_key 加密存储(掩码/空 = 不设置)
if data.get('api_key') and data['api_key'] != MASK:
data['api_key'] = encrypt_api_key(data['api_key'])
else:
data.pop('api_key', None)
# 端点下标校验JSON 或逗号分隔均可)
data['endpoint_ids'] = _norm_endpoint_ids(data.get('endpoint_ids', ''))
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
recs = await sor.sqlExe("SELECT endpoints FROM llm_vendor WHERE id=${v}$", {"v": vid})
await sor.sqlExe("COMMIT", {})
if not recs:
raise ValueError('供应商不存在')
eps = []
try:
eps = json.loads(getattr(recs[0], 'endpoints', '') or '') or []
except Exception:
eps = []
for i in json.loads(data['endpoint_ids'] or '[]'):
if int(i) < 0 or int(i) >= len(eps):
raise ValueError('端点下标 %s 超出供应商端点目录(共 %d 个端点)' % (i, len(eps)))
data['id'] = getID()
async with db.sqlorContext(dbname) as sor:
await sor.C('llm_account', data)
result['success'] = True
result['message'] = '创建成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def update_llm_account(params_kw):
"""账号更新api_key 非空才重新加密(空=不改)。"""
result = {'success': False, 'message': ''}
try:
data = _clean(params_kw)
if data.get('api_key') and data['api_key'] != MASK:
data['api_key'] = encrypt_api_key(data['api_key'])
else:
data.pop('api_key', None)
if 'endpoint_ids' in data:
data['endpoint_ids'] = _norm_endpoint_ids(data.get('endpoint_ids', ''))
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
await sor.U('llm_account', data)
result['success'] = True
result['message'] = '更新成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def delete_llm_account(params_kw):
result = {'success': False, 'message': ''}
try:
aid = (params_kw or {}).get('id', '')
if not aid:
raise ValueError('缺少 id')
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
recs = await sor.sqlExe(
"SELECT COUNT(*) AS c FROM llm_model WHERE account_id=${a}$", {"a": aid})
await sor.sqlExe("COMMIT", {})
cnt = int(getattr(recs[0], 'c', 0)) if recs else 0
if cnt > 0:
raise ValueError('该账号被 %d 个模型设为默认账号,请先解除后再删除' % cnt)
await sor.sqlExe("DELETE FROM llm_account WHERE id=${a}$", {"a": aid})
await sor.sqlExe("COMMIT", {})
result['success'] = True
result['message'] = '删除成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def create_llm_model(params_kw):
result = {'success': False, 'message': ''}
try:
data = _clean(params_kw)
name = (data.get('name', '') or '').strip()
if not name:
raise ValueError('模型注册名不能为空')
if not data.get('vendor_id'):
raise ValueError('必须选择供应商')
dp = data.get('default_params', '') or ''
if isinstance(dp, (list, dict)):
dp = json.dumps(dp, ensure_ascii=False)
if dp.strip():
json.loads(dp) # 校验合法
data['default_params'] = dp
# 后续步骤模板链JSON 数组校验
qp = data.get('query_profile_ids', '') or ''
if isinstance(qp, (list, dict)):
qp = json.dumps(qp, ensure_ascii=False)
if qp.strip():
ql = json.loads(qp)
if not isinstance(ql, list):
raise ValueError('query_profile_ids 必须是 JSON 数组')
data['query_profile_ids'] = qp
data['id'] = getID()
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
recs = await sor.sqlExe("SELECT id FROM llm_model WHERE name=${n}$", {"n": name})
await sor.sqlExe("COMMIT", {})
if recs:
raise ValueError('模型注册名「%s」已存在(平台内唯一)' % name)
await sor.C('llm_model', data)
result['success'] = True
result['message'] = '创建成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def update_llm_model(params_kw):
result = {'success': False, 'message': ''}
try:
data = _clean(params_kw)
if 'default_params' in data:
dp = data.get('default_params', '') or ''
if isinstance(dp, (list, dict)):
dp = json.dumps(dp, ensure_ascii=False)
if dp.strip():
json.loads(dp)
data['default_params'] = dp
if 'query_profile_ids' in data:
qp = data.get('query_profile_ids', '') or ''
if isinstance(qp, (list, dict)):
qp = json.dumps(qp, ensure_ascii=False)
if qp.strip():
ql = json.loads(qp)
if not isinstance(ql, list):
raise ValueError('query_profile_ids 必须是 JSON 数组')
data['query_profile_ids'] = qp
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
await sor.U('llm_model', data)
result['success'] = True
result['message'] = '更新成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def delete_llm_model(params_kw):
"""模型删除前置检查:被策略引用时禁止。"""
result = {'success': False, 'message': ''}
try:
mid = (params_kw or {}).get('id', '')
if not mid:
raise ValueError('缺少 id')
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
recs = await sor.sqlExe(
"SELECT COUNT(*) AS c FROM llm_org_policy WHERE primary_model_id=${m}$", {"m": mid})
await sor.sqlExe("COMMIT", {})
cnt = int(getattr(recs[0], 'c', 0)) if recs else 0
if cnt > 0:
raise ValueError('该模型被 %d 个组织策略设为主模型,请先调整策略' % cnt)
await sor.sqlExe("DELETE FROM llm_model WHERE id=${m}$", {"m": mid})
await sor.sqlExe("COMMIT", {})
result['success'] = True
result['message'] = '删除成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def update_llm_org_policy(params_kw):
"""组织策略保存:校验模型存在、备链去重。"""
result = {'success': False, 'message': ''}
try:
data = _clean(params_kw)
backups = data.get('backup_model_ids', '') or ''
if isinstance(backups, (list, dict)):
backups = json.dumps(backups, ensure_ascii=False)
bl = []
if backups.strip():
bl = json.loads(backups)
if not isinstance(bl, list):
raise ValueError('backup_model_ids 必须是 JSON 数组')
seen = []
for b in bl:
if b and b not in seen:
seen.append(b)
bl = seen
data['backup_model_ids'] = json.dumps(bl, ensure_ascii=False) if bl else ''
primary = data.get('primary_model_id', '') or ''
if primary and primary in bl:
raise ValueError('主模型不能同时出现在备模型链中')
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
if primary:
recs = await sor.sqlExe("SELECT id FROM llm_model WHERE id=${m}$", {"m": primary})
await sor.sqlExe("COMMIT", {})
if not recs:
raise ValueError('主模型不存在或已停用')
org_id = data.get('org_id', '')
if org_id:
recs = await sor.sqlExe("SELECT id FROM llm_org_policy WHERE org_id=${o}$", {"o": org_id})
await sor.sqlExe("COMMIT", {})
if recs:
data['id'] = getattr(recs[0], 'id', '')
await sor.U('llm_org_policy', data)
else:
data['id'] = getID()
await sor.C('llm_org_policy', data)
else:
await sor.U('llm_org_policy', data)
result['success'] = True
result['message'] = '保存成功'
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def llm_usage_query(params_kw):
"""用量查询(列表页工具栏):按机构/模型/账号/状态/时间范围过滤,返回汇总+明细。"""
result = {'success': False, 'message': ''}
try:
pk = params_kw or {}
conds = ["1=1"]
params = {}
if pk.get('org_id'):
conds.append("org_id=${o}$"); params['o'] = pk['org_id']
if pk.get('model_id'):
conds.append("model_id=${m}$"); params['m'] = pk['model_id']
if pk.get('account_id'):
conds.append("account_id=${a}$"); params['a'] = pk['account_id']
if pk.get('status'):
conds.append("status=${s}$"); params['s'] = pk['status']
if pk.get('date_from'):
conds.append("created_at>=${df}$"); params['df'] = pk['date_from']
if pk.get('date_to'):
conds.append("created_at<=${dt}$"); params['dt'] = pk['date_to'] + ' 23:59:59'
where = " AND ".join(conds)
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
recs = await sor.sqlExe(
"SELECT u.id, u.org_id, u.user_id, u.model_id, m.name AS model_name, "
"u.account_id, a.name AS account_name, u.endpoint_region, "
"u.req_tokens, u.resp_tokens, u.cost, u.charge, u.ppid, u.task_ref, "
"u.status, u.note, u.created_at "
"FROM llm_usage u "
"LEFT JOIN llm_model m ON m.id=u.model_id "
"LEFT JOIN llm_account a ON a.id=u.account_id "
"WHERE %s ORDER BY u.created_at DESC LIMIT 200" % where, params)
sums = await sor.sqlExe(
"SELECT COUNT(*) AS cnt, COALESCE(SUM(req_tokens),0) AS req_t, "
"COALESCE(SUM(resp_tokens),0) AS resp_t, COALESCE(SUM(cost),0) AS cost_s, "
"COALESCE(SUM(charge),0) AS charge_s FROM llm_usage WHERE %s" % where, params)
await sor.sqlExe("COMMIT", {})
rows = []
for r in (recs or []):
d = {}
for k, v in vars(r).items():
if not callable(v):
d[k] = v
rows.append(d)
s = sums[0] if sums else None
result['success'] = True
result['rows'] = rows
result['summary'] = {
'count': int(getattr(s, 'cnt', 0)) if s else 0,
'req_tokens': int(getattr(s, 'req_t', 0)) if s else 0,
'resp_tokens': int(getattr(s, 'resp_t', 0)) if s else 0,
'cost': round(float(getattr(s, 'cost_s', 0)), 6) if s else 0,
'charge': round(float(getattr(s, 'charge_s', 0)), 6) if s else 0,
}
except Exception as e:
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
async def llm_dashboard(params_kw):
"""治理总览(首页卡片数据源):各机构策略/余额/近期用量。"""
result = {'success': True}
try:
db, dbname = _get_sor()
async with db.sqlorContext(dbname) as sor:
vendors = await sor.sqlExe("SELECT COUNT(*) AS c FROM llm_vendor WHERE status='active'", {})
accounts = await sor.sqlExe(
"SELECT COUNT(*) AS c, COALESCE(SUM(balance),0) AS b FROM llm_account WHERE status='active'", {})
models = await sor.sqlExe("SELECT COUNT(*) AS c FROM llm_model WHERE status='active'", {})
orgs = await sor.sqlExe("SELECT COUNT(*) AS c FROM llm_org_policy WHERE status='active'", {})
usage = await sor.sqlExe(
"SELECT COUNT(*) AS c, COALESCE(SUM(charge),0) AS ch FROM llm_usage WHERE status='ok'", {})
await sor.sqlExe("COMMIT", {})
result['vendors'] = int(getattr(vendors[0], 'c', 0)) if vendors else 0
result['accounts'] = int(getattr(accounts[0], 'c', 0)) if accounts else 0
result['account_balance'] = round(float(getattr(accounts[0], 'b', 0)), 4) if accounts else 0
result['models'] = int(getattr(models[0], 'c', 0)) if models else 0
result['org_policies'] = int(getattr(orgs[0], 'c', 0)) if orgs else 0
result['usage_calls'] = int(getattr(usage[0], 'c', 0)) if usage else 0
result['usage_charge'] = round(float(getattr(usage[0], 'ch', 0)), 4) if usage else 0
except Exception as e:
result['success'] = False
result['message'] = str(e)
return json.dumps(result, ensure_ascii=False, default=str)
def load_llm():
"""模块加载:注册全部函数到 ServerEnv。"""
load_gateway()
env = ServerEnv()
env.create_llm_vendor = create_llm_vendor
env.update_llm_vendor = update_llm_vendor
env.delete_llm_vendor = delete_llm_vendor
env.create_llm_account = create_llm_account
env.update_llm_account = update_llm_account
env.delete_llm_account = delete_llm_account
env.create_llm_model = create_llm_model
env.update_llm_model = update_llm_model
env.delete_llm_model = delete_llm_model
env.update_llm_org_policy = update_llm_org_policy
env.llm_usage_query = llm_usage_query
env.llm_dashboard = llm_dashboard
logger.info("[pipeline_llm] v1.0.0 loaded — 模型治理模块就绪")