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

@@ -42,7 +42,7 @@ from api.db.services.compilation_template_group_service import CompilationTempla
from api.db.joint_services.memory_message_service import handle_save_to_memory_task
from api.db.joint_services.tenant_model_service import (
get_tenant_default_model_by_type,
get_model_config_from_provider_instance,
resolve_model_config,
get_model_config_by_id,
)
from api.db.services.llm_service import LLMBundle
@@ -321,11 +321,17 @@ class TaskHandler:
try:
if ctx.tenant_embd_id:
try:
embd_model_config = get_model_config_by_id(task_tenant_id, ctx.tenant_embd_id)
embd_model_config = get_model_config_by_id(
task_tenant_id, LLMType.EMBEDDING, ctx.tenant_embd_id
)
except LookupError:
embd_model_config = get_model_config_from_provider_instance(task_tenant_id, LLMType.EMBEDDING, task_embedding_id)
embd_model_config = resolve_model_config(
task_tenant_id, LLMType.EMBEDDING, task_embedding_id
)
elif task_embedding_id:
embd_model_config = get_model_config_from_provider_instance(task_tenant_id, LLMType.EMBEDDING, task_embedding_id)
embd_model_config = resolve_model_config(
task_tenant_id, LLMType.EMBEDDING, task_embedding_id
)
else:
embd_model_config = get_tenant_default_model_by_type(task_tenant_id, LLMType.EMBEDDING)
embedding_model = LLMBundle(task_tenant_id, embd_model_config, lang=task_language)
@@ -381,7 +387,7 @@ class TaskHandler:
return
# Bind LLM for raptor
chat_model_config = get_model_config_from_provider_instance(task_tenant_id, LLMType.CHAT, kb_task_llm_id)
chat_model_config = resolve_model_config(task_tenant_id, LLMType.CHAT, kb_task_llm_id)
with LLMBundle(task_tenant_id, chat_model_config, lang=ctx.language) as chat_model:
# Run RAPTOR
raptor_service = RaptorService(ctx=ctx)
@@ -497,7 +503,7 @@ class TaskHandler:
graphrag_conf = kb_parser_config.get("graphrag", {})
start_ts = timer()
chat_model_config = get_model_config_from_provider_instance(task_tenant_id, LLMType.CHAT, kb_task_llm_id)
chat_model_config = resolve_model_config(task_tenant_id, LLMType.CHAT, kb_task_llm_id)
with LLMBundle(task_tenant_id, chat_model_config, lang=task_language) as chat_model:
with_resolution = graphrag_conf.get("resolution", False)
with_community = graphrag_conf.get("community", False)
@@ -780,7 +786,7 @@ class TaskHandler:
def _build_toc(cls, ctx: TaskContext, docs: List[Dict], progress_cb: Callable) -> Optional[Dict]:
"""Build table of contents."""
progress_cb(msg="Start to generate table of content ...")
chat_model_config = get_model_config_from_provider_instance(ctx.tenant_id, LLMType.CHAT, ctx.llm_id)
chat_model_config = resolve_model_config(ctx.tenant_id, LLMType.CHAT, ctx.llm_id)
with LLMBundle(ctx.tenant_id, chat_model_config, lang=ctx.language) as chat_mdl:
docs = sorted(
docs,