Fix: display embedding model name in list dataset api response (#17832)

This commit is contained in:
Lynn
2026-08-05 14:51:07 +08:00
committed by GitHub
parent 510e1197c6
commit 02923bfc46
4 changed files with 41 additions and 3 deletions

View File

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

View File

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

View File

@@ -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"},
]

View File

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