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

@@ -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"])

View File

@@ -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)

View File

@@ -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: