refa: resolve tenant model refs consistently (#16744)

This commit is contained in:
buua436
2026-07-09 14:02:08 +08:00
committed by GitHub
parent 794fcc2517
commit 6a77523bf0
51 changed files with 300 additions and 225 deletions

View File

@@ -37,7 +37,7 @@ from api.db.services.user_service import TenantService
from common.metadata_utils import apply_meta_data_filter
from api.db.services.search_service import SearchService
from api.db.services.user_service import UserTenantService
from api.db.joint_services.tenant_model_service import get_tenant_default_model_by_type, get_model_config_from_provider_instance
from api.db.joint_services.tenant_model_service import get_tenant_default_model_by_type, resolve_model_config
from common.misc_utils import thread_pool_exec
from api.utils.api_utils import get_error_data_result, get_json_result, add_tenant_id_to_kwargs, get_result, get_request_json, server_error_response, validate_request
from rag.app.tag import label_question
@@ -405,7 +405,7 @@ async def retrieval_test_embedded(tenant_id=None):
if meta_data_filter.get("method") in ["auto", "semi_auto"]:
chat_id = search_config.get("chat_id", "")
if chat_id:
chat_model_config = await thread_pool_exec(get_model_config_from_provider_instance, tenant_id, LLMType.CHAT, chat_id)
chat_model_config = await thread_pool_exec(resolve_model_config, tenant_id, LLMType.CHAT, chat_id)
else:
chat_model_config = await thread_pool_exec(get_tenant_default_model_by_type, tenant_id, LLMType.CHAT)
chat_mdl = LLMBundle(tenant_id, chat_model_config)
@@ -450,12 +450,12 @@ async def retrieval_test_embedded(tenant_id=None):
if langs:
_question = await cross_languages(kb.tenant_id, None, _question, langs)
embd_model_config = await thread_pool_exec(get_model_config_from_provider_instance, kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
embd_model_config = await thread_pool_exec(resolve_model_config, kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
embd_mdl = LLMBundle(kb.tenant_id, embd_model_config)
rerank_mdl = None
if rerank_id:
rerank_model_config = await thread_pool_exec(get_model_config_from_provider_instance, tenant_id, LLMType.RERANK, rerank_id)
rerank_model_config = await thread_pool_exec(resolve_model_config, tenant_id, LLMType.RERANK, rerank_id)
rerank_mdl = LLMBundle(kb.tenant_id, rerank_model_config)
if req.get("keyword", False):
@@ -523,7 +523,7 @@ async def related_questions_embedded(tenant_id=None):
chat_id = search_config.get("chat_id", "")
if chat_id:
chat_model_config = await thread_pool_exec(get_model_config_from_provider_instance, tenant_id, LLMType.CHAT, chat_id)
chat_model_config = await thread_pool_exec(resolve_model_config, tenant_id, LLMType.CHAT, chat_id)
else:
chat_model_config = await thread_pool_exec(get_tenant_default_model_by_type, tenant_id, LLMType.CHAT)
chat_mdl = LLMBundle(tenant_id, chat_model_config)

View File

@@ -27,7 +27,7 @@ from quart import Response, request
from api.apps import current_user, login_required
from api.apps.restful_apis._generation_params import merge_generation_config, pop_generation_config
from api.db.joint_services.tenant_model_service import get_tenant_default_model_by_type, get_model_config_from_provider_instance, get_api_key
from api.db.joint_services.tenant_model_service import get_api_key, get_tenant_default_model_by_type, resolve_model_config
from api.db.services.chunk_feedback_service import ChunkFeedbackService
from api.db.services.conversation_service import ConversationService, structure_answer
from api.db.services.dialog_service import DialogService, async_chat, gen_mindmap
@@ -277,10 +277,10 @@ async def _validate_llm_id(llm_id, tenant_id, llm_setting=None):
model_type = "chat"
try:
await thread_pool_exec(
get_model_config_from_provider_instance,
resolve_model_config,
tenant_id=tenant_id,
model_name=llm_id,
model_type=model_type,
model_ref=llm_id,
)
except Exception as e:
logging.error(f"Fail to get model config for {llm_id}: {e}")
@@ -298,7 +298,7 @@ async def _validate_rerank_id(rerank_id, tenant_id):
return None
try:
await thread_pool_exec(
get_model_config_from_provider_instance,
resolve_model_config,
tenant_id=tenant_id,
model_name=rerank_id,
model_type="rerank",
@@ -1127,7 +1127,7 @@ async def recommendation():
chat_id = search_config.get("chat_id", "")
if chat_id:
chat_model_config = get_model_config_from_provider_instance(current_user.id, LLMType.CHAT, chat_id)
chat_model_config = resolve_model_config(current_user.id, LLMType.CHAT, chat_id)
else:
chat_model_config = get_tenant_default_model_by_type(current_user.id, LLMType.CHAT)
chat_mdl = LLMBundle(current_user.id, chat_model_config)

View File

@@ -27,7 +27,7 @@ from quart import request
from api.apps import login_required
from api.db.joint_services.tenant_model_service import (
split_model_name,
get_model_config_from_provider_instance,
resolve_model_config,
get_tenant_default_model_by_type,
)
from api.db.db_models import Document, Task
@@ -374,12 +374,12 @@ async def retrieval_test(tenant_id):
e, kb = KnowledgebaseService.get_by_id(kb_ids[0])
if not e:
return get_error_data_result(message="Dataset not found!")
embd_model_config = get_model_config_from_provider_instance(kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
embd_model_config = resolve_model_config(kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
embd_mdl = LLMBundle(kb.tenant_id, embd_model_config)
rerank_mdl = None
if req.get("rerank_id"):
rerank_model_config = get_model_config_from_provider_instance(kb.tenant_id, LLMType.RERANK, req["rerank_id"])
rerank_model_config = resolve_model_config(kb.tenant_id, LLMType.RERANK, req["rerank_id"])
rerank_mdl = LLMBundle(kb.tenant_id, rerank_model_config)
if langs:
@@ -902,7 +902,7 @@ async def add_chunk(tenant_id, dataset_id, document_id):
d["doc_type_kwd"] = "image"
embd_id = DocumentService.get_embd_id(document_id)
model_config = get_model_config_from_provider_instance(dataset_tenant_id, LLMType.EMBEDDING.value, embd_id)
model_config = resolve_model_config(dataset_tenant_id, LLMType.EMBEDDING.value, embd_id)
embd_mdl = TenantLLMService.model_instance(model_config)
v, c = embd_mdl.encode([doc.name, req["content"] if not d["question_kwd"] else "\n".join(d["question_kwd"])])
v = 0.1 * v[0] + 0.9 * v[1]
@@ -1048,7 +1048,7 @@ async def update_chunk(tenant_id, dataset_id, document_id, chunk_id):
d["doc_type_kwd"] = "image"
embd_id = DocumentService.get_embd_id(document_id)
model_config = get_model_config_from_provider_instance(dataset_tenant_id, LLMType.EMBEDDING.value, embd_id)
model_config = resolve_model_config(dataset_tenant_id, LLMType.EMBEDDING.value, embd_id)
embd_mdl = TenantLLMService.model_instance(model_config)
if doc.parser_id == ParserType.QA:
arr = [t for t in re.split(r"[\n\t]", d["content_with_weight"]) if len(t) > 1]

View File

@@ -27,7 +27,7 @@ from api.db.services.document_service import DocumentService
from api.db.services.doc_metadata_service import DocMetadataService
from api.db.services.knowledgebase_service import KnowledgebaseService
from api.db.services.llm_service import LLMBundle
from api.db.joint_services.tenant_model_service import get_tenant_default_model_by_type, get_model_config_from_provider_instance
from api.db.joint_services.tenant_model_service import get_tenant_default_model_by_type, resolve_model_config
from common.metadata_utils import meta_filter, convert_conditions
from api.apps import login_required
from api.utils.api_utils import add_tenant_id_to_kwargs, build_error_result, get_request_json, get_json_result
@@ -261,7 +261,7 @@ async def retrieval(tenant_id):
kb_id,
)
return build_error_result(message="No authorization.", code=RetCode.AUTHENTICATION_ERROR)
model_config = get_model_config_from_provider_instance(kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
model_config = resolve_model_config(kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
embd_mdl = LLMBundle(kb.tenant_id, model_config)
if metadata_condition:
doc_ids.extend(meta_filter(metas, convert_conditions(metadata_condition), metadata_condition.get("logic", "and")))

View File

@@ -19,6 +19,7 @@ from quart import request
from api.apps import login_required
from api.apps.services import models_api_service
from api.db.services.user_service import TenantService
from api.utils.api_utils import (
add_tenant_id_to_kwargs,
get_error_argument_result,
@@ -39,6 +40,16 @@ def get_added_models(tenant_id: str):
security:
- ApiKeyAuth: []
parameters:
- in: query
name: owner_tenant_id
type: string
required: false
description: "If provided, list models from the owner tenant's scope after access validation."
- in: query
name: type
type: string
required: false
description: "Model type filter (chat, embedding, rerank, asr, vision, tts, ocr)."
- in: header
name: Authorization
type: string
@@ -70,8 +81,18 @@ def get_added_models(tenant_id: str):
type: boolean
"""
model_type_filter = request.args.get("type")
owner_tenant_id = request.args.get("owner_tenant_id")
try:
success, result = models_api_service.list_tenant_added_models(tenant_id, model_type_filter)
target_tenant_id = tenant_id
if owner_tenant_id:
if owner_tenant_id != tenant_id:
joined_tenants = TenantService.get_joined_tenants_by_user_id(tenant_id)
allowed_tenant_ids = {tenant_id, *(tenant["tenant_id"] for tenant in joined_tenants)}
if owner_tenant_id not in allowed_tenant_ids:
return get_error_data_result(message="Permission denied")
target_tenant_id = owner_tenant_id
success, result = models_api_service.list_tenant_added_models(target_tenant_id, model_type_filter)
if success:
return get_result(data=result)
else:

View File

@@ -23,7 +23,7 @@ from api.apps import current_user, login_required
from api.apps.restful_apis._generation_params import extract_generation_config, merge_generation_config
from api.db.services.dialog_service import DialogService, async_chat
from api.db.services.doc_metadata_service import DocMetadataService
from api.db.joint_services.tenant_model_service import get_model_config_from_provider_instance, get_api_key
from api.db.joint_services.tenant_model_service import resolve_model_config, get_api_key
from api.utils.api_utils import get_error_data_result, get_request_json, validate_request
from common.constants import RetCode, StatusEnum
from common.metadata_utils import convert_conditions, meta_filter
@@ -40,9 +40,9 @@ def _validate_llm_id(llm_id, tenant_id, llm_setting=None):
model_type = "chat"
try:
get_model_config_from_provider_instance(
resolve_model_config(
tenant_id=tenant_id,
model_name=llm_id,
model_ref=llm_id,
model_type=model_type,
)
except Exception as e:

View File

@@ -18,8 +18,8 @@ import json
import os
import re
from api.db.joint_services.tenant_model_service import get_model_config_from_provider_instance
from common.constants import PAGERANK_FLD
from api.db.joint_services.tenant_model_service import resolve_model_config, resolve_model_id
from common.constants import PAGERANK_FLD, LLMType
from common import settings
from api.db.db_models import File
from api.db.services.document_service import DocumentService, queue_raptor_o_graphrag_tasks
@@ -28,6 +28,7 @@ from api.db.services.file_service import FileService
from api.db.services.knowledgebase_service import KnowledgebaseService, validate_dataset_embedding_models
from api.db.services.connector_service import Connector2KbService
from api.db.services.task_service import GRAPH_RAPTOR_FAKE_DOC_ID, TaskService
from api.db.services.tenant_model_service import TenantModelService
from api.db.services.user_service import TenantService, UserService, UserTenantService
from common.constants import FileSource, StatusEnum
from api.utils.api_utils import deep_merge, get_parser_config, remap_dictionary_keys, verify_embedding_availability
@@ -328,6 +329,11 @@ async def update_dataset(tenant_id: str, dataset_id: str, req: dict):
ok, err = verify_embedding_availability(req["embd_id"], tenant_id)
if not ok:
return False, err
ok, _ = TenantModelService.get_by_id(req["embd_id"])
if ok:
req["tenant_embd_id"] = req["embd_id"]
else:
req["tenant_embd_id"] = resolve_model_id(tenant_id, LLMType.EMBEDDING, req["embd_id"])
if "pagerank" in req and req["pagerank"] != kb.pagerank:
if os.environ.get("DOC_ENGINE", "elasticsearch") == "infinity":
@@ -1011,7 +1017,7 @@ async def search(dataset_id: str, tenant_id: str, req: dict):
if meta_data_filter.get("method") in ["auto", "semi_auto"]:
chat_id = search_config.get("chat_id", "")
if chat_id:
chat_model_config = get_model_config_from_provider_instance(tenant_id, LLMType.CHAT, search_config["chat_id"])
chat_model_config = resolve_model_config(tenant_id, LLMType.CHAT, search_config["chat_id"])
else:
chat_model_config = get_tenant_default_model_by_type(tenant_id, LLMType.CHAT)
chat_mdl = LLMBundle(tenant_id, chat_model_config)
@@ -1045,15 +1051,15 @@ async def search(dataset_id: str, tenant_id: str, req: dict):
if langs:
_question = await cross_languages(kb.tenant_id, None, _question, langs)
if kb.embd_id:
embd_model_config = get_model_config_from_provider_instance(kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
embd_model_config = resolve_model_config(kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
else:
embd_model_config = get_tenant_default_model_by_type(kb.tenant_id, LLMType.EMBEDDING)
embd_mdl = LLMBundle(kb.tenant_id, embd_model_config)
rerank_mdl = None
rerank_id = search_config.get("rerank_id") or req.get("rerank_id")
rerank_id = req.get("rerank_id") or search_config.get("rerank_id")
if rerank_id:
rerank_model_config = get_model_config_from_provider_instance(kb.tenant_id, LLMType.RERANK.value, rerank_id)
rerank_model_config = resolve_model_config(kb.tenant_id, LLMType.RERANK.value, rerank_id)
rerank_mdl = LLMBundle(kb.tenant_id, rerank_model_config)
if search_config.get("keyword", req.get("keyword", False)):
@@ -1243,7 +1249,7 @@ def check_embedding(dataset_id: str, tenant_id: str, req: dict):
if not ok:
return False, err
embd_model_config = get_model_config_from_provider_instance(kb.tenant_id, LLMType.EMBEDDING, embd_id)
embd_model_config = resolve_model_config(kb.tenant_id, LLMType.EMBEDDING, embd_id)
emb_mdl = LLMBundle(kb.tenant_id, embd_model_config)
n = int(req.get("check_num", 5))
@@ -1400,7 +1406,7 @@ async def search_datasets(tenant_id: str, req: dict):
if meta_data_filter.get("method") in ["auto", "semi_auto"]:
chat_id = search_config.get("chat_id", "")
if chat_id:
chat_model_config = get_model_config_from_provider_instance(tenant_id, LLMType.CHAT, search_config["chat_id"])
chat_model_config = resolve_model_config(tenant_id, LLMType.CHAT, search_config["chat_id"])
else:
chat_model_config = get_tenant_default_model_by_type(tenant_id, LLMType.CHAT)
chat_mdl = LLMBundle(tenant_id, chat_model_config)
@@ -1438,13 +1444,15 @@ async def search_datasets(tenant_id: str, req: dict):
embd_mdl = None
if kb.embd_id:
embd_model_config = get_model_config_from_provider_instance(kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
embd_mdl = LLMBundle(kb.tenant_id, embd_model_config)
embd_model_config = resolve_model_config(kb.tenant_id, LLMType.EMBEDDING, kb.embd_id)
else:
embd_model_config = get_tenant_default_model_by_type(kb.tenant_id, LLMType.EMBEDDING)
embd_mdl = LLMBundle(kb.tenant_id, embd_model_config)
rerank_mdl = None
rerank_id = search_config.get("rerank_id") or req.get("rerank_id")
rerank_id = req.get("rerank_id") or search_config.get("rerank_id")
if rerank_id:
rerank_model_config = get_model_config_from_provider_instance(kb.tenant_id, LLMType.RERANK.value, rerank_id)
rerank_model_config = resolve_model_config(kb.tenant_id, LLMType.RERANK.value, rerank_id)
rerank_mdl = LLMBundle(kb.tenant_id, rerank_model_config)
if search_config.get("keyword", req.get("keyword", False)):

View File

@@ -285,14 +285,17 @@ def list_tenant_added_models(tenant_id: str, model_type_filter: str = None):
model_records = TenantModelService.get_models_by_provider_ids_and_instance_ids(provider_ids, list({instance.id for instance in instances}))
target_type_records = [record for record in model_records if record.model_type & model_type_filter_bin] if model_type_filter_bin else model_records
factory_rank_mapping = {factory["name"]: -_to_int(factory.get("rank", "500")) for factory in FACTORY_LLM_INFOS}
model_rank_map: dict = {}
for factory in FACTORY_LLM_INFOS:
for llm in factory.get("llm", []):
model_rank_map[(factory["name"], llm["llm_name"])] = _to_int(llm.get("rank", 500))
factory_rank_mapping = {factory["name"]: -_to_int(factory.get("rank", "500")) for factory in FACTORY_LLM_INFOS}
added_models = [
{
"model_id": model_record.id,
"tenant_id": provider_info_map[model_record.provider_id].tenant_id,
"tenant_name": tenant.name,
"model_type": get_model_type_human(model_record.model_type),
"name": model_record.model_name,
"provider_id": model_record.provider_id,
@@ -313,6 +316,9 @@ def list_tenant_added_models(tenant_id: str, model_type_filter: str = None):
if not tei_already_added:
added_models.append(
{
"model_id": "",
"tenant_id": tenant.id,
"tenant_name": tenant.name,
"model_type": ["embedding"],
"name": tei_model,
"provider_id": "",

View File

@@ -20,7 +20,7 @@ import asyncio
from common.constants import LLMType, ActiveStatusEnum, ModelVerifyStatusEnum
from common.settings import FACTORY_LLM_INFOS
from api.db.joint_services.tenant_model_service import get_model_config_from_provider_instance, delete_models_by_instance_ids, delete_instances_by_provider_ids
from api.db.joint_services.tenant_model_service import resolve_model_config, delete_models_by_instance_ids, delete_instances_by_provider_ids
from api.db.services.tenant_model_provider_service import TenantModelProviderService
from api.db.services.tenant_model_instance_service import TenantModelInstanceService
from api.db.services.tenant_model_service import TenantModelService
@@ -1141,7 +1141,7 @@ async def chat_to_model(tenant_id: str, provider_id_or_name: str, instance_id_or
# Get model config
composite_name = f"{model_name}@{instance_name}@{provider_name}"
try:
model_config = get_model_config_from_provider_instance(tenant_id, LLMType.CHAT, composite_name)
model_config = resolve_model_config(tenant_id, LLMType.CHAT, composite_name)
except LookupError:
return False, f"Model '{composite_name}' not authorized"