Handle searching dataset without embedding model (#16742)

### Summary

Handle searching dataset without embedding model

In this PR, Searching datasets with different embedding models or
searching dataset with/without embedding models are not allowed. We will
improve the behavior later.
This commit is contained in:
qinling0210
2026-07-09 11:38:55 +08:00
committed by GitHub
parent 1430d0e431
commit ae96e636e9
14 changed files with 142 additions and 107 deletions

View File

@@ -331,7 +331,7 @@ type InsertChunksFromFileRequest struct {
// @Security ApiKeyAuth
// @Param request body InsertChunksFromFileRequest true "insert chunks request"
// @Success 200 {object} map[string]interface{}
// @Router /v1/tenant/insert_chunks_from_file [post]
// @Router /v1/tenant/dev_insert_chunks_from_file [post]
func (h *TenantHandler) InsertChunksFromFile(c *gin.Context) {
_, errorCode, errorMessage := GetUser(c)
if errorCode != common.CodeSuccess {
@@ -409,7 +409,7 @@ type InsertMetadataFromFileRequest struct {
// @Security ApiKeyAuth
// @Param request body InsertMetadataFromFileRequest true "insert metadata request"
// @Success 200 {object} map[string]interface{}
// @Router /v1/tenant/insert_metadata_from_file [post]
// @Router /v1/tenant/dev_insert_metadata_from_file [post]
func (h *TenantHandler) InsertMetadataFromFile(c *gin.Context) {
user, errorCode, errorMessage := GetUser(c)
if errorCode != common.CodeSuccess {

View File

@@ -275,12 +275,12 @@ func (r *Router) Setup(engine *gin.Engine) {
tenant := v1.Group("/tenant")
{
tenant.GET("/list", r.tenantHandler.TenantList)
tenant.POST("/chunk_store", r.tenantHandler.CreateChunkStore) // Internal API only for GO
tenant.DELETE("/chunk_store", r.tenantHandler.DeleteChunkStore) // Internal API only for GO
tenant.POST("/metadata_store", r.tenantHandler.CreateMetadataStore) // Internal API only for GO
tenant.DELETE("/metadata_store", r.tenantHandler.DeleteMetadataStore) // Internal API only for GO
tenant.POST("/insert_chunks_from_file", r.tenantHandler.InsertChunksFromFile) // Internal API only for GO
tenant.POST("/insert_metadata_from_file", r.tenantHandler.InsertMetadataFromFile) // Internal API only for GO
tenant.POST("/chunk_store", r.tenantHandler.CreateChunkStore) // Internal API only for GO
tenant.DELETE("/chunk_store", r.tenantHandler.DeleteChunkStore) // Internal API only for GO
tenant.POST("/metadata_store", r.tenantHandler.CreateMetadataStore) // Internal API only for GO
tenant.DELETE("/metadata_store", r.tenantHandler.DeleteMetadataStore) // Internal API only for GO
tenant.POST("/dev_insert_chunks_from_file", r.tenantHandler.InsertChunksFromFile) // Internal API only for GO
tenant.POST("/dev_insert_metadata_from_file", r.tenantHandler.InsertMetadataFromFile) // Internal API only for GO
}
// Document routes

View File

@@ -289,12 +289,8 @@ func (s *ChatService) validateCreateDatasetIDs(value interface{}, tenantID strin
kbs = append(kbs, kb)
}
embedIDs := make(map[string]struct{}, len(kbs))
for _, kb := range kbs {
embedIDs[s.splitModelNameAndFactory(kb.EmbdID)] = struct{}{}
}
if len(embedIDs) > 1 {
return nil, fmt.Errorf("Datasets use different embedding models: %v", getEmbdIDs(kbs))
if err := validateDatasetEmbeddingModels(kbs); err != nil {
return nil, err
}
return normalizedIDs, nil
}
@@ -645,24 +641,21 @@ const (
pyDefaultEmptyResponse = "Sorry! No relevant content was found in the knowledge base!"
)
// splitModelNameAndFactory extracts the base model name (removes vendor suffix)
// splitModelNameAndFactory extracts the base model name by stripping
// provider and instance suffixes, matching Python's rsplit("@", 2)[0].
func (s *ChatService) splitModelNameAndFactory(embdID string) string {
// Remove vendor suffix (e.g., "model@openai" -> "model")
if idx := strings.LastIndex(embdID, "@"); idx > 0 {
return embdID[:idx]
// Strip the provider segment.
base := embdID[:idx]
// Strip the instance segment (second-to-last @).
if idx2 := strings.LastIndex(base, "@"); idx2 > 0 {
return base[:idx2]
}
return base
}
return embdID
}
// getEmbdIDs extracts embedding IDs from knowledge bases
func getEmbdIDs(kbs []*entity.Knowledgebase) []string {
ids := make([]string, len(kbs))
for i, kb := range kbs {
ids[i] = kb.EmbdID
}
return ids
}
func (s *ChatService) getOwnedValidChat(userID, chatID string) (*entity.Chat, error) {
chat, err := s.chatDAO.GetByIDAndStatus(chatID, string(entity.StatusValid))
if err != nil {

View File

@@ -245,7 +245,14 @@ func (s *ChatPipelineService) AsyncChat(
// === Phase 4: Bind Models (embedding, rerank, chat, TTS) + ToolCall ===
common.Info("Phase 4: Bind Models (embedding, rerank, chat, TTS)")
timer.Enter(common.PhaseBindModels)
kbs, embModel, rerankModel, chatModel, ttsModel := s.getModels(ctx, chat)
kbs, embModel, rerankModel, chatModel, ttsModel, err := s.getModels(ctx, chat)
if err != nil {
out <- AsyncChatResult{
Answer: fmt.Sprintf("**ERROR**: %s", err.Error()),
Final: true,
}
return
}
// Toolcall binding
if toolcallSession, hasSession := kwargs["toolcall_session"]; hasSession && toolcallSession != nil {
@@ -1918,6 +1925,7 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat)
*modelModule.RerankModel,
*modelModule.ChatModel,
*modelModule.ChatModel, // TTS model
error,
) {
kbDAO := dao.NewKnowledgebaseDAO()
@@ -1942,27 +1950,22 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat)
// Embedding model.
var embModel *modelModule.EmbeddingModel
if len(kbs) > 0 {
// All KBs must share the same embedding model.
embdIDs := make(map[string]bool)
for _, kb := range kbs {
if kb.EmbdID != "" {
embdIDs[kb.EmbdID] = true
}
if err := validateDatasetEmbeddingModels(kbs); err != nil {
return nil, nil, nil, nil, nil, err
}
if len(embdIDs) > 1 {
// Multiple embedding models across KBs — error.
common.Warn("Knowledge bases use different embedding models")
}
if len(embdIDs) == 1 {
for embdID := range embdIDs {
embdTenantID := kbs[0].TenantID
driver, modelName, apiConfig, maxTokens, err := s.ModelProviderSvc.GetModelConfigFromProviderInstance(
embdTenantID, entity.ModelTypeEmbedding, embdID,
)
if err == nil {
embModel = modelModule.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens)
}
if kbs[0].EmbdID != "" {
embdTenantID := kbs[0].TenantID
driver, modelName, apiConfig, maxTokens, err := s.ModelProviderSvc.GetModelConfigFromProviderInstance(
embdTenantID, entity.ModelTypeEmbedding, kbs[0].EmbdID,
)
if err != nil {
common.Warn("Failed to get embedding model for chat retrieval",
zap.String("embdID", kbs[0].EmbdID),
zap.String("tenantID", embdTenantID),
zap.Error(err))
return nil, nil, nil, nil, nil, fmt.Errorf("failed to get embedding model: %w", err)
}
embModel = modelModule.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens)
}
}
@@ -1997,7 +2000,7 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat)
}
}
return kbs, embModel, rerankModel, chatModel, ttsModel
return kbs, embModel, rerankModel, chatModel, ttsModel, nil
}
// lastUserQuestion returns the content of the most recent user message in

View File

@@ -102,6 +102,47 @@ const (
graphPhaseCommunityDone = "community_done"
)
// validateDatasetEmbeddingModels checks that all given datasets use the same
// embedding model (or all use none). Returns an error on mismatch.
func validateDatasetEmbeddingModels(kbs []*entity.Knowledgebase) error {
embdIDs := make(map[string]struct{})
hasEmbd := false
noEmbd := false
for _, kb := range kbs {
if kb.EmbdID != "" {
hasEmbd = true
baseName := kb.EmbdID
if idx := strings.LastIndex(kb.EmbdID, "@"); idx > 0 {
baseName = kb.EmbdID[:idx]
// Strip the second-to-last @-segment too (instance name),
// matching Python's _base_model_name which uses rsplit("@", 2).
if idx2 := strings.LastIndex(baseName, "@"); idx2 > 0 {
baseName = baseName[:idx2]
}
}
embdIDs[baseName] = struct{}{}
} else {
noEmbd = true
}
}
if hasEmbd && noEmbd {
return fmt.Errorf("Cannot search across datasets where some have embedding models and others do not.")
}
if len(embdIDs) > 1 {
return fmt.Errorf("Datasets use different embedding models: %v", getEmbdIDs(kbs))
}
return nil
}
// getEmbdIDs extracts embedding IDs from knowledge bases.
func getEmbdIDs(kbs []*entity.Knowledgebase) []string {
ids := make([]string, len(kbs))
for i, kb := range kbs {
ids[i] = kb.EmbdID
}
return ids
}
// DatasetService implements the RESTful dataset APIs from dataset_api.py.
type DatasetService struct {
kbDAO *dao.KnowledgebaseDAO
@@ -1872,13 +1913,8 @@ func (d *DatasetService) SearchDatasets(req *SearchDatasetsRequest, userID strin
}
// Check if all kbs have the same embedding model
if len(kbRecords) > 1 {
firstEmbdID := kbRecords[0].EmbdID
for i := 1; i < len(kbRecords); i++ {
if kbRecords[i].EmbdID != firstEmbdID {
return nil, fmt.Errorf("Datasets use different embedding models.")
}
}
if err := validateDatasetEmbeddingModels(kbRecords); err != nil {
return nil, err
}
// Override request fields with values from saved search config (if search_id is provided)
@@ -2050,36 +2086,23 @@ func (d *DatasetService) SearchDatasets(req *SearchDatasetsRequest, userID strin
return nil, fmt.Errorf("failed to get embedding model by embd_id: %w", embErr)
}
embeddingModel = models.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens)
} else {
driver, modelName, apiConfig, maxTokens, err := modelProviderSvc.GetTenantDefaultModelByType(tenantIDs[0], entity.ModelTypeEmbedding)
if err != nil {
return nil, fmt.Errorf("failed to get tenant default embedding model: %w", err)
}
embeddingModel = models.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens)
common.Info("Fetched embedding model for retrieval",
zap.String("tenantID", tenantIDs[0]),
zap.String("modelName", modelName))
}
modelNameStr := ""
if embeddingModel.ModelName != nil {
modelNameStr = *embeddingModel.ModelName
}
common.Info("Fetched embedding model for retrieval",
zap.String("tenantID", tenantIDs[0]),
zap.String("modelName", modelNameStr))
// Get rerank model if rerankID is specified
var rerankModel *models.RerankModel
if rerankID != "" {
driver, modelName, apiConfig, _, rErr := modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeRerank, rerankID)
if rErr != nil {
return nil, fmt.Errorf("failed to get rerank model by rerank_id: %w", rErr)
}
rerankModel = models.NewRerankModel(driver, &modelName, apiConfig)
}
if rerankModel != nil {
common.Info("Fetched rerank model",
zap.String("tenantID", tenantIDs[0]),
zap.String("modelName", *rerankModel.ModelName))
zap.String("modelName", modelName))
}
retrievalReq := &nlp.RetrievalRequest{