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