mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-22 16:36:47 +08:00
refa: resolve tenant model refs consistently (#16744)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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")))
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)):
|
||||
|
||||
@@ -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": "",
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user