mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-07 08:01:13 +08:00
Fix: display embedding model name in list dataset api response (#17832)
This commit is contained in:
@@ -18,7 +18,7 @@ import json
|
||||
import os
|
||||
import re
|
||||
|
||||
from api.db.joint_services.tenant_model_service import resolve_model_config, resolve_model_id
|
||||
from api.db.joint_services.tenant_model_service import resolve_model_config, resolve_model_id, get_composite_model_name_by_ids
|
||||
from common.constants import PAGERANK_FLD, LLMType
|
||||
from common import settings
|
||||
from api.db.db_models import File
|
||||
@@ -478,6 +478,10 @@ def list_datasets(tenant_id: str, args: dict):
|
||||
if status_by_kb:
|
||||
kb["parsing_status"] = status_by_kb.get(kb["id"], {})
|
||||
response_data_list.append(remap_dictionary_keys(kb))
|
||||
|
||||
embed_model_names = get_composite_model_name_by_ids([m["embedding_model"] for m in response_data_list])
|
||||
for response_data in response_data_list:
|
||||
response_data["embedding_model_name"] = embed_model_names.get(response_data["embedding_model"], "")
|
||||
return True, {"data": response_data_list, "total": total}
|
||||
|
||||
|
||||
|
||||
@@ -551,6 +551,38 @@ def get_composite_model_name_by_id(model_id: str) -> str:
|
||||
return f"{model_obj.model_name}@{instance_obj.instance_name}@{provider_obj.provider_name}"
|
||||
|
||||
|
||||
def get_composite_model_name_by_ids(model_ids: list[str]) -> dict[str, str]:
|
||||
"""Convert a list of tenant_model.id values to a dict mapping each id
|
||||
to its composite model name string ``model_name@instance_name@provider_name``.
|
||||
Model ids that cannot be resolved are silently skipped.
|
||||
"""
|
||||
if not model_ids:
|
||||
return {}
|
||||
|
||||
models = list(TenantModelService.get_by_ids(model_ids))
|
||||
if not models:
|
||||
return {}
|
||||
|
||||
instance_ids = list({m.instance_id for m in models})
|
||||
provider_ids = list({m.provider_id for m in models})
|
||||
|
||||
instances = list(TenantModelInstanceService.get_by_ids(instance_ids)) if instance_ids else []
|
||||
instance_map = {i.id: i for i in instances}
|
||||
|
||||
providers = list(TenantModelProviderService.get_by_ids(provider_ids)) if provider_ids else []
|
||||
provider_map = {p.id: p for p in providers}
|
||||
|
||||
result: dict[str, str] = {}
|
||||
for m in models:
|
||||
inst = instance_map.get(m.instance_id)
|
||||
prov = provider_map.get(m.provider_id)
|
||||
if not inst or not prov:
|
||||
continue
|
||||
result[m.id] = f"{m.model_name}@{inst.instance_name}@{prov.provider_name}"
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def ensure_mistral_ocr_from_env(tenant_id: str) -> str | None:
|
||||
return _ensure_ocr_provider_from_env(
|
||||
tenant_id,
|
||||
|
||||
@@ -82,6 +82,7 @@ def _load_list_datasets_module(monkeypatch, *, kbs, parsing_status_by_kb):
|
||||
_stub(
|
||||
monkeypatch,
|
||||
"api.db.joint_services.tenant_model_service",
|
||||
get_composite_model_name_by_ids=MagicMock(),
|
||||
resolve_model_config=MagicMock(),
|
||||
resolve_model_id=MagicMock(),
|
||||
)
|
||||
@@ -188,8 +189,8 @@ def _load_list_datasets_module(monkeypatch, *, kbs, parsing_status_by_kb):
|
||||
|
||||
def _stub_kbs():
|
||||
return [
|
||||
{"id": "kb-a", "tenant_id": "tenant-1", "name": "Alpha"},
|
||||
{"id": "kb-b", "tenant_id": "tenant-1", "name": "Beta"},
|
||||
{"id": "kb-a", "tenant_id": "tenant-1", "name": "Alpha", "embedding_model": "emb-a"},
|
||||
{"id": "kb-b", "tenant_id": "tenant-1", "name": "Beta", "embedding_model": "emb-b"},
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -113,6 +113,7 @@ def _load_delete_datasets_module(monkeypatch, *, f2d_rows, file_filter_delete):
|
||||
_stub(
|
||||
monkeypatch,
|
||||
"api.db.joint_services.tenant_model_service",
|
||||
get_composite_model_name_by_ids=MagicMock(),
|
||||
get_model_config_from_provider_instance=MagicMock(),
|
||||
resolve_model_config=MagicMock(),
|
||||
resolve_model_id=MagicMock(),
|
||||
|
||||
Reference in New Issue
Block a user