mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-22 16:36:47 +08:00
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:
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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{
|
||||
|
||||
Reference in New Issue
Block a user