diff --git a/pipeline_llm/gateway.py b/pipeline_llm/gateway.py index f55e9a7..7d2ad33 100644 --- a/pipeline_llm/gateway.py +++ b/pipeline_llm/gateway.py @@ -199,7 +199,7 @@ async def _model_chain(sor, org_id: str, model_name: str, purpose: str = ''): for mid in chain_ids: recs = await sor.sqlExe( "SELECT id, name, vendor_id, vendor_model_id, capability, default_params, " - "ppid, org_id, profile_id, sync_mode, query_profile_ids " + "ppid, org_id, profile_id, sync_mode, query_profile_ids, account_id " "FROM llm_model WHERE id=${i}$ AND status='active' LIMIT 1", {"i": mid}) await sor.sqlExe("COMMIT", {}) if recs: @@ -223,10 +223,20 @@ async def _vendor_endpoints(sor, vendor_id: str): async def _account_candidates(sor, model: dict): - """模型的候选账号池:同供应商、启用中。""" + """模型的候选账号池:同供应商、启用中。 + + 模型配置了 account_id 时(如某模型只允许用某个账号充值/密钥), + 候选池收窄为该账号——模型级账号绑定优先于供应商级轮转。 + """ + where = "vendor_id=${v}$ AND status='active'" + params = {"v": model.get('vendor_id', '')} + bound = (model.get('account_id') or '').strip() + if bound: + where += " AND id=${a}$" + params['a'] = bound recs = await sor.sqlExe( "SELECT id, name, api_key, endpoint_ids, balance, status FROM llm_account " - "WHERE vendor_id=${v}$ AND status='active'", {"v": model.get('vendor_id', '')}) + "WHERE " + where, params) await sor.sqlExe("COMMIT", {}) return [_row_to_dict(r) for r in (recs or [])]