mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-28 11:48:10 +08:00
refa: resolve tenant model refs consistently (#16744)
This commit is contained in:
@@ -27,7 +27,7 @@ from api.db.db_models import Task
|
||||
from api.db.services.task_service import TaskService
|
||||
from api.db.services.memory_service import MemoryService
|
||||
from api.db.services.llm_service import LLMBundle
|
||||
from api.db.joint_services.tenant_model_service import get_model_config_from_provider_instance, get_model_config_by_id
|
||||
from api.db.joint_services.tenant_model_service import resolve_model_config, get_model_config_by_id
|
||||
from api.utils.memory_utils import get_memory_type_human
|
||||
from memory.services.messages import MessageService
|
||||
from memory.services.query import MsgTextQuery, get_vector
|
||||
@@ -169,11 +169,11 @@ async def extract_by_llm(tenant_id: str, tenant_llm_id: str | None, extract_conf
|
||||
user_prompts.append({"role": "user", "content": PromptAssembler.assemble_user_prompt(conversation_content, conversation_time, conversation_time)})
|
||||
if tenant_llm_id:
|
||||
try:
|
||||
llm_config = get_model_config_by_id(tenant_id, tenant_llm_id)
|
||||
llm_config = get_model_config_by_id(tenant_id, LLMType.CHAT, tenant_llm_id)
|
||||
except LookupError:
|
||||
llm_config = get_model_config_from_provider_instance(tenant_id, LLMType.CHAT, llm_id)
|
||||
llm_config = resolve_model_config(tenant_id, LLMType.CHAT, llm_id)
|
||||
else:
|
||||
llm_config = get_model_config_from_provider_instance(tenant_id, LLMType.CHAT, llm_id)
|
||||
llm_config = resolve_model_config(tenant_id, LLMType.CHAT, llm_id)
|
||||
with LLMBundle(tenant_id, llm_config) as llm:
|
||||
if task_id:
|
||||
TaskService.update_progress(task_id, {"progress": 0.15, "progress_msg": timestamp_to_date(current_timestamp()) + " " + "Prepared prompts and LLM."})
|
||||
@@ -196,11 +196,11 @@ async def extract_by_llm(tenant_id: str, tenant_llm_id: str | None, extract_conf
|
||||
async def embed_and_save(memory, message_list: list[dict], task_id: str=None):
|
||||
if memory.tenant_embd_id:
|
||||
try:
|
||||
embd_model_config = get_model_config_by_id(memory.tenant_id, memory.tenant_embd_id)
|
||||
embd_model_config = get_model_config_by_id(memory.tenant_id, LLMType.EMBEDDING, memory.tenant_embd_id)
|
||||
except LookupError:
|
||||
embd_model_config = get_model_config_from_provider_instance(memory.tenant_id, LLMType.EMBEDDING, memory.embd_id)
|
||||
embd_model_config = resolve_model_config(memory.tenant_id, LLMType.EMBEDDING, memory.embd_id)
|
||||
else:
|
||||
embd_model_config = get_model_config_from_provider_instance(memory.tenant_id, LLMType.EMBEDDING, memory.embd_id)
|
||||
embd_model_config = resolve_model_config(memory.tenant_id, LLMType.EMBEDDING, memory.embd_id)
|
||||
with LLMBundle(memory.tenant_id, embd_model_config) as embedding_model:
|
||||
if task_id:
|
||||
TaskService.update_progress(task_id, {"progress": 0.65, "progress_msg": timestamp_to_date(current_timestamp()) + " " + "Prepared embedding model."})
|
||||
@@ -270,7 +270,7 @@ def query_message(filter_dict: dict, params: dict):
|
||||
question = params["query"]
|
||||
question = question.strip()
|
||||
memory = memory_list[0]
|
||||
embd_model_config = get_model_config_from_provider_instance(memory.tenant_id, LLMType.EMBEDDING, memory.embd_id)
|
||||
embd_model_config = resolve_model_config(memory.tenant_id, LLMType.EMBEDDING, memory.embd_id)
|
||||
embd_model = LLMBundle(memory.tenant_id, embd_model_config)
|
||||
match_dense = get_vector(question, embd_model, similarity=params["similarity_threshold"])
|
||||
match_text, _ = MsgTextQuery().question(question, min_match=params["similarity_threshold"])
|
||||
|
||||
@@ -195,10 +195,10 @@ def get_tenant_default_model_by_type(tenant_id: str, model_type: str | enum.Enum
|
||||
# Prefer resolving by tenant_model.id when available
|
||||
if model_id:
|
||||
try:
|
||||
return get_model_config_by_id(tenant_id, model_id)
|
||||
return get_model_config_by_id(tenant_id, model_type, model_id)
|
||||
except LookupError:
|
||||
logger.warning("tenant_model id=%s not found, falling back to model_name lookup for %s", model_id, model_name)
|
||||
return get_model_config_from_provider_instance(tenant_id, model_type, model_name)
|
||||
return resolve_model_config(tenant_id, model_type, model_name)
|
||||
|
||||
|
||||
def split_model_name(model_name: str):
|
||||
@@ -249,6 +249,11 @@ def _resolve_instance_for_model(provider_obj, instance_name: str, model_name: st
|
||||
|
||||
raise LookupError(f"Instance {instance_name} not found for model {model_name}.")
|
||||
|
||||
def resolve_model_config(tenant_id, model_type: str | enum.Enum, model_ref: str):
|
||||
try:
|
||||
return get_model_config_by_id(tenant_id, model_type, model_ref)
|
||||
except LookupError:
|
||||
return get_model_config_from_provider_instance(tenant_id, model_type, model_ref)
|
||||
|
||||
def get_model_config_from_provider_instance(tenant_id, model_type: str | enum.Enum, model_name: str):
|
||||
pure_model_name, instance_name, provider_name = split_model_name(model_name)
|
||||
@@ -312,16 +317,22 @@ def get_model_config_from_provider_instance(tenant_id, model_type: str | enum.En
|
||||
raise LookupError(f"Model {model_name} not found for model {model_type_val}")
|
||||
|
||||
|
||||
def get_model_config_by_id(tenant_id: str, model_id: str):
|
||||
def get_model_config_by_id(tenant_id: str, model_type: str | enum.Enum, model_id: str):
|
||||
"""Get model config from tenant_model by its id (CharField PK)."""
|
||||
model_type_val = model_type if isinstance(model_type, str) else model_type.value
|
||||
model_type_bin = calculate_model_type(model_type_val)
|
||||
exist, model_obj = TenantModelService.get_by_id(model_id)
|
||||
if not exist:
|
||||
raise LookupError(f"TenantModel id={model_id} not found.")
|
||||
if model_obj.status != ActiveStatusEnum.ACTIVE.value:
|
||||
if model_obj.status == ActiveStatusEnum.INACTIVE.value:
|
||||
raise LookupError(f"TenantModel id={model_id} is disabled.")
|
||||
if model_obj.status == ActiveStatusEnum.UNSUPPORTED.value:
|
||||
raise LookupError(f"TenantModel id={model_id} cannot be used as {model_type_val} model.")
|
||||
if not (model_obj.model_type & model_type_bin):
|
||||
raise LookupError(f"TenantModel id={model_id} cannot be used as {model_type_val} model.")
|
||||
|
||||
provider_obj = TenantModelProviderService.get_by_id(model_obj.provider_id)
|
||||
if not provider_obj:
|
||||
ok, provider_obj = TenantModelProviderService.get_by_id(model_obj.provider_id)
|
||||
if not ok:
|
||||
raise LookupError(f"Provider id={model_obj.provider_id} not found for model id={model_id}.")
|
||||
|
||||
# Validate that tenant_id owns the provider or is a joined tenant of the provider's owner.
|
||||
@@ -331,8 +342,8 @@ def get_model_config_by_id(tenant_id: str, model_id: str):
|
||||
if provider_obj.tenant_id not in joined_tenant_ids:
|
||||
raise LookupError(f"Tenant {tenant_id} has no access to provider owned by tenant {provider_obj.tenant_id}.")
|
||||
|
||||
instance_obj = TenantModelInstanceService.get_by_id(model_obj.instance_id)
|
||||
if not instance_obj:
|
||||
ok, instance_obj = TenantModelInstanceService.get_by_id(model_obj.instance_id)
|
||||
if not ok:
|
||||
raise LookupError(f"Instance id={model_obj.instance_id} not found for model id={model_id}.")
|
||||
|
||||
api_key, is_tool, api_key_payload = _decode_api_key_config(instance_obj.api_key)
|
||||
@@ -344,7 +355,7 @@ def get_model_config_by_id(tenant_id: str, model_id: str):
|
||||
"api_key": api_key,
|
||||
"llm_name": model_obj.model_name,
|
||||
"api_base": extra_fields.get("base_url", ""),
|
||||
"model_type": model_obj.model_type,
|
||||
"model_type": model_type_val,
|
||||
"is_tools": model_extra.get("is_tools", is_tool),
|
||||
"max_tokens": model_extra.get("max_tokens") or 8192,
|
||||
}
|
||||
@@ -432,6 +443,17 @@ def get_api_key(tenant_id: str, model_name: str):
|
||||
instance_obj = _resolve_instance_for_model(provider_obj, instance_name, model_name)
|
||||
return instance_obj.api_key
|
||||
|
||||
def get_model_type_by_id(model_id: str):
|
||||
exist, model_obj = TenantModelService.get_by_id(model_id)
|
||||
if not exist:
|
||||
raise LookupError(f"TenantModel id={model_id} not found.")
|
||||
return get_model_type_human(model_obj.model_type)
|
||||
|
||||
def resolve_model_type(tenant_id: str, model_ref: str):
|
||||
try:
|
||||
return get_model_type_by_id(model_ref)
|
||||
except LookupError:
|
||||
return get_model_type_by_name(tenant_id, model_ref)
|
||||
|
||||
def get_model_type_by_name(tenant_id: str, model_name: str):
|
||||
pure_model_name, instance_name, provider_name = split_model_name(model_name)
|
||||
|
||||
@@ -39,7 +39,7 @@ from api.utils.reference_metadata_utils import (
|
||||
enrich_chunks_with_document_metadata,
|
||||
resolve_reference_metadata_preferences,
|
||||
)
|
||||
from api.db.joint_services.tenant_model_service import get_tenant_default_model_by_type, get_model_config_from_provider_instance, get_model_type_by_name, get_model_config_by_id
|
||||
from api.db.joint_services.tenant_model_service import get_tenant_default_model_by_type, resolve_model_config, resolve_model_type, get_model_config_by_id
|
||||
from common.time_utils import current_timestamp, datetime_format
|
||||
from common.text_utils import normalize_arabic_digits
|
||||
from rag.advanced_rag.knowlege_compile.mind_map_extractor import MindMapExtractor
|
||||
@@ -295,23 +295,23 @@ async def async_chat_solo(dialog, messages, stream=True, session_id=None):
|
||||
if dialog.llm_id:
|
||||
if dialog.tenant_llm_id:
|
||||
try:
|
||||
llm_types = get_model_type_by_name(dialog.tenant_id, dialog.llm_id)
|
||||
llm_types = resolve_model_type(dialog.tenant_id, dialog.llm_id)
|
||||
if "chat" in llm_types:
|
||||
model_config = get_model_config_by_id(dialog.tenant_id, dialog.tenant_llm_id)
|
||||
model_config = get_model_config_by_id(dialog.tenant_id, LLMType.CHAT, dialog.tenant_llm_id)
|
||||
else:
|
||||
model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
model_config = resolve_model_config(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
except LookupError:
|
||||
llm_types = get_model_type_by_name(dialog.tenant_id, dialog.llm_id)
|
||||
llm_types = resolve_model_type(dialog.tenant_id, dialog.llm_id)
|
||||
if "chat" in llm_types:
|
||||
model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
model_config = resolve_model_config(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
else:
|
||||
model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
model_config = resolve_model_config(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
else:
|
||||
llm_types = get_model_type_by_name(dialog.tenant_id, dialog.llm_id)
|
||||
llm_types = resolve_model_type(dialog.tenant_id, dialog.llm_id)
|
||||
if "chat" in llm_types:
|
||||
model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
model_config = resolve_model_config(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
else:
|
||||
model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
model_config = resolve_model_config(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
else:
|
||||
model_config = get_tenant_default_model_by_type(dialog.tenant_id, LLMType.CHAT)
|
||||
|
||||
@@ -364,7 +364,7 @@ def get_models(dialog, trace_context=None, langfuse_session_id=None):
|
||||
|
||||
if kbs and kbs[0].embd_id:
|
||||
embd_owner_tenant_id = kbs[0].tenant_id
|
||||
embd_model_config = get_model_config_from_provider_instance(embd_owner_tenant_id, LLMType.EMBEDDING, kbs[0].embd_id)
|
||||
embd_model_config = resolve_model_config(embd_owner_tenant_id, LLMType.EMBEDDING, kbs[0].embd_id)
|
||||
embd_mdl = LLMBundle(embd_owner_tenant_id, embd_model_config, trace_context=trace_context, langfuse_session_id=langfuse_session_id)
|
||||
if not embd_mdl:
|
||||
raise LookupError("Embedding model(%s) not found" % kbs[0].embd_id)
|
||||
@@ -372,11 +372,11 @@ def get_models(dialog, trace_context=None, langfuse_session_id=None):
|
||||
if dialog.llm_id:
|
||||
if dialog.tenant_llm_id:
|
||||
try:
|
||||
chat_model_config = get_model_config_by_id(dialog.tenant_id, dialog.tenant_llm_id)
|
||||
chat_model_config = get_model_config_by_id(dialog.tenant_id, LLMType.CHAT, dialog.tenant_llm_id)
|
||||
except LookupError:
|
||||
chat_model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
chat_model_config = resolve_model_config(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
else:
|
||||
chat_model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
chat_model_config = resolve_model_config(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
else:
|
||||
chat_model_config = get_tenant_default_model_by_type(dialog.tenant_id, LLMType.CHAT)
|
||||
|
||||
@@ -385,11 +385,11 @@ def get_models(dialog, trace_context=None, langfuse_session_id=None):
|
||||
if dialog.rerank_id:
|
||||
if dialog.tenant_rerank_id:
|
||||
try:
|
||||
rerank_model_config = get_model_config_by_id(dialog.tenant_id, dialog.tenant_rerank_id)
|
||||
rerank_model_config = get_model_config_by_id(dialog.tenant_id, LLMType.RERANK, dialog.tenant_rerank_id)
|
||||
except LookupError:
|
||||
rerank_model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.RERANK, dialog.rerank_id)
|
||||
rerank_model_config = resolve_model_config(dialog.tenant_id, LLMType.RERANK, dialog.rerank_id)
|
||||
else:
|
||||
rerank_model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.RERANK, dialog.rerank_id)
|
||||
rerank_model_config = resolve_model_config(dialog.tenant_id, LLMType.RERANK, dialog.rerank_id)
|
||||
rerank_mdl = LLMBundle(dialog.tenant_id, rerank_model_config, trace_context=trace_context, langfuse_session_id=langfuse_session_id)
|
||||
|
||||
if dialog.prompt_config.get("tts"):
|
||||
@@ -583,23 +583,23 @@ async def async_chat(dialog, messages, stream=True, **kwargs):
|
||||
if dialog.llm_id:
|
||||
if dialog.tenant_llm_id:
|
||||
try:
|
||||
llm_types = get_model_type_by_name(dialog.tenant_id, dialog.llm_id)
|
||||
llm_types = resolve_model_type(dialog.tenant_id, dialog.llm_id)
|
||||
if "chat" in llm_types:
|
||||
llm_model_config = get_model_config_by_id(dialog.tenant_id, dialog.tenant_llm_id)
|
||||
llm_model_config = get_model_config_by_id(dialog.tenant_id, LLMType.CHAT, dialog.tenant_llm_id)
|
||||
else:
|
||||
llm_model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
llm_model_config = resolve_model_config(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
except LookupError:
|
||||
llm_types = get_model_type_by_name(dialog.tenant_id, dialog.llm_id)
|
||||
llm_types = resolve_model_type(dialog.tenant_id, dialog.llm_id)
|
||||
if "chat" in llm_types:
|
||||
llm_model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
llm_model_config = resolve_model_config(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
else:
|
||||
llm_model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
llm_model_config = resolve_model_config(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
else:
|
||||
llm_types = get_model_type_by_name(dialog.tenant_id, dialog.llm_id)
|
||||
llm_types = resolve_model_type(dialog.tenant_id, dialog.llm_id)
|
||||
if "chat" in llm_types:
|
||||
llm_model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
llm_model_config = resolve_model_config(dialog.tenant_id, LLMType.CHAT, dialog.llm_id)
|
||||
else:
|
||||
llm_model_config = get_model_config_from_provider_instance(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
llm_model_config = resolve_model_config(dialog.tenant_id, LLMType.IMAGE2TEXT, dialog.llm_id)
|
||||
else:
|
||||
llm_model_config = get_tenant_default_model_by_type(dialog.tenant_id, LLMType.CHAT)
|
||||
|
||||
@@ -1683,12 +1683,12 @@ async def async_ask(question, kb_ids, tenant_id, chat_llm_name=None, search_conf
|
||||
is_knowledge_graph = all([kb.parser_id == ParserType.KG for kb in kbs])
|
||||
retriever = settings.retriever if not is_knowledge_graph else settings.kg_retriever
|
||||
embd_owner_tenant_id = kbs[0].tenant_id
|
||||
embd_model_config = get_model_config_from_provider_instance(embd_owner_tenant_id, LLMType.EMBEDDING, embedding_list[0])
|
||||
embd_model_config = resolve_model_config(embd_owner_tenant_id, LLMType.EMBEDDING, embedding_list[0])
|
||||
embd_mdl = LLMBundle(embd_owner_tenant_id, embd_model_config)
|
||||
chat_model_config = get_model_config_from_provider_instance(tenant_id, LLMType.CHAT, chat_llm_name)
|
||||
chat_model_config = resolve_model_config(tenant_id, LLMType.CHAT, chat_llm_name)
|
||||
chat_mdl = LLMBundle(tenant_id, chat_model_config)
|
||||
if rerank_id:
|
||||
rerank_model_config = get_model_config_from_provider_instance(tenant_id, LLMType.RERANK, rerank_id)
|
||||
rerank_model_config = resolve_model_config(tenant_id, LLMType.RERANK, rerank_id)
|
||||
rerank_mdl = LLMBundle(tenant_id, rerank_model_config)
|
||||
max_tokens = chat_mdl.max_length
|
||||
tenant_ids = list(set([kb.tenant_id for kb in kbs]))
|
||||
@@ -1794,16 +1794,16 @@ async def gen_mindmap(question, kb_ids, tenant_id, search_config={}):
|
||||
return {"error": "No KB selected"}
|
||||
tenant_ids = list(set([kb.tenant_id for kb in kbs]))
|
||||
embd_owner_tenant_id = kbs[0].tenant_id
|
||||
embd_model_config = get_model_config_from_provider_instance(embd_owner_tenant_id, LLMType.EMBEDDING, kbs[0].embd_id)
|
||||
embd_model_config = resolve_model_config(embd_owner_tenant_id, LLMType.EMBEDDING, kbs[0].embd_id)
|
||||
embd_mdl = LLMBundle(embd_owner_tenant_id, embd_model_config)
|
||||
chat_id = search_config.get("chat_id", "")
|
||||
if chat_id:
|
||||
chat_model_config = get_model_config_from_provider_instance(tenant_id, LLMType.CHAT, chat_id)
|
||||
chat_model_config = resolve_model_config(tenant_id, LLMType.CHAT, chat_id)
|
||||
else:
|
||||
chat_model_config = get_tenant_default_model_by_type(tenant_id, LLMType.CHAT)
|
||||
chat_mdl = LLMBundle(tenant_id, chat_model_config)
|
||||
if rerank_id:
|
||||
rerank_model_config = get_model_config_from_provider_instance(tenant_id, LLMType.RERANK, rerank_id)
|
||||
rerank_model_config = resolve_model_config(tenant_id, LLMType.RERANK, rerank_id)
|
||||
rerank_mdl = LLMBundle(tenant_id, rerank_model_config)
|
||||
|
||||
if meta_data_filter:
|
||||
|
||||
Reference in New Issue
Block a user