Files
ragflow/internal/agent/tool/retrieval_nlp_test.go
Hz_ a55a438b42 fix(go-agent): finish retrieval component (#17845)
- Prefer canonical `dataset_ids` over legacy `kb_ids` in Canvas
retrieval components.
- Cover selected and explicitly cleared dataset IDs with regression
tests.
- Add cross-language retrieval with dataset-specific embedding and
rerank models.
- Support vector and keyword similarity controls, metadata filtering,
TOC enhancement, and child-chunk expansion.
- Route Canvas retrieval across datasets and memories with compatible
embedding validation.
2026-08-05 15:48:54 +08:00

514 lines
15 KiB
Go

//
// Copyright 2026 The InfiniFlow Authors. All Rights Reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//
// retrieval_nlp_test.go — NLPRetrievalAdapter tests.
//
// The adapter's Search method calls into the real
// nlp.RetrievalService, which requires a doc engine + document
// DAO. Those require network access and aren't worth standing up
// in a unit test. Instead we exercise the translation layer
// (translateChunk + helpers) directly — that's the surface that
// differs between the nlp chunk shape and the agent-tool chunk
// shape, and it's the only piece the adapter owns.
package tool
import (
"context"
"errors"
"math"
"testing"
"ragflow/internal/entity"
modelModule "ragflow/internal/entity/models"
"ragflow/internal/service/nlp"
"gorm.io/gorm"
)
// floatEqual compares two floats with a small epsilon so
// arithmetic like 0.4+0.8/2 == 0.6000000000000001 doesn't fail.
func floatEqual(a, b float64) bool { return math.Abs(a-b) < 1e-9 }
// TestTranslateChunk_FullFields: every chunk field is present and
// maps to the expected RetrievalChunk field.
func TestTranslateChunk_FullFields(t *testing.T) {
raw := map[string]any{
"chunk_id": "ck-42",
"doc_id": "doc-7",
"docnm_kwd": "report.pdf",
"kb_id": "kb-1",
"image_id": "img-1",
"position_int": [][]float64{{1, 2, 3, 4}},
"content_with_weight": "the answer is 42",
"content_ltks": "answer 42",
"similarity": 0.87,
"term_similarity": 0.5,
"vector_similarity": 0.9,
}
got := translateChunk(raw)
if got.ID != "ck-42" {
t.Errorf("ID = %q, want \"ck-42\"", got.ID)
}
if got.Content != "the answer is 42" {
t.Errorf("Content = %q, want content_with_weight", got.Content)
}
if got.DocumentID != "doc-7" {
t.Errorf("DocumentID = %q, want \"doc-7\"", got.DocumentID)
}
if got.DocumentName != "report.pdf" {
t.Errorf("DocumentName = %q, want \"report.pdf\"", got.DocumentName)
}
if got.DatasetID != "kb-1" {
t.Errorf("DatasetID = %q, want \"kb-1\"", got.DatasetID)
}
if got.ImageID != "img-1" {
t.Errorf("ImageID = %q, want \"img-1\"", got.ImageID)
}
if got.Positions == nil {
t.Errorf("Positions is nil, want position_int payload")
}
if got.Score != 0.87 {
t.Errorf("Score = %v, want 0.87 (similarity preferred)", got.Score)
}
if got.TermSimilarity != 0.5 {
t.Errorf("TermSimilarity = %v, want 0.5", got.TermSimilarity)
}
if got.VectorSimilarity != 0.9 {
t.Errorf("VectorSimilarity = %v, want 0.9", got.VectorSimilarity)
}
}
// TestTranslateChunk_ContentFallback: when content_with_weight is
// empty, content_ltks is used.
func TestTranslateChunk_ContentFallback(t *testing.T) {
raw := map[string]any{
"chunk_id": "ck-1",
"content_ltks": "answer 42",
"doc_id": "doc-1",
"similarity": 0.5,
}
got := translateChunk(raw)
if got.Content != "answer 42" {
t.Errorf("Content = %q, want content_ltks fallback", got.Content)
}
}
// TestTranslateChunk_EmptyContent: both content fields empty →
// empty string (don't crash, don't synthesise).
func TestTranslateChunk_EmptyContent(t *testing.T) {
raw := map[string]any{
"chunk_id": "ck-1",
"doc_id": "doc-1",
"similarity": 0.5,
}
got := translateChunk(raw)
if got.Content != "" {
t.Errorf("Content = %q, want \"\"", got.Content)
}
}
// TestTranslateChunk_ScoreFallback: when "similarity" is missing,
// the average of term_similarity + vector_similarity is used.
func TestTranslateChunk_ScoreFallback(t *testing.T) {
raw := map[string]any{
"chunk_id": "ck-1",
"term_similarity": 0.4,
"vector_similarity": 0.8,
}
got := translateChunk(raw)
if !floatEqual(got.Score, 0.6) {
t.Errorf("Score = %v, want ~0.6 (avg of term+vec)", got.Score)
}
}
// TestTranslateChunk_ScoreOnlyOneSub: only one of term_similarity
// or vector_similarity present → that one wins (not averaged
// against zero, which would be misleading).
func TestTranslateChunk_ScoreOnlyOneSub(t *testing.T) {
cases := []struct {
name string
raw map[string]any
want float64
}{
{"term only", map[string]any{"term_similarity": 0.7}, 0.7},
{"vector only", map[string]any{"vector_similarity": 0.3}, 0.3},
}
for _, tc := range cases {
got := translateChunk(tc.raw)
if !floatEqual(got.Score, tc.want) {
t.Errorf("%s: Score = %v, want %v", tc.name, got.Score, tc.want)
}
}
}
// TestTranslateChunk_NumberTypesTolerated: numeric similarity
// fields may come back as float64, float32, int, or int64 depending
// on how the upstream serialised them. All should coerce.
func TestTranslateChunk_NumberTypesTolerated(t *testing.T) {
cases := []struct {
name string
val any
want float64
}{
{"float64", float64(0.5), 0.5},
{"float32", float32(0.5), 0.5},
{"int", int(1), 1},
{"int64", int64(1), 1},
}
for _, tc := range cases {
raw := map[string]any{"similarity": tc.val}
got := translateChunk(raw)
if !floatEqual(got.Score, tc.want) {
t.Errorf("%s: Score = %v, want %v", tc.name, got.Score, tc.want)
}
}
}
// TestTranslateChunk_WrongTypesIgnored: a similarity field that's
// a string or nil must not crash. We fall back to the sub-scores
// (which are also missing in this test → zero).
func TestTranslateChunk_WrongTypesIgnored(t *testing.T) {
cases := []map[string]any{
{"similarity": "0.5"},
{"similarity": nil},
{"similarity": []any{0.5}},
}
for _, raw := range cases {
got := translateChunk(raw)
if got.Score != 0 {
t.Errorf("wrong-type similarity: Score = %v, want 0", got.Score)
}
}
}
// TestTranslateChunk_WrongStringTypesIgnored: a string field
// that's actually a number must not crash.
func TestTranslateChunk_WrongStringTypesIgnored(t *testing.T) {
raw := map[string]any{
"chunk_id": 42, // int, not string
"doc_id": nil, // nil
}
got := translateChunk(raw)
if got.ID != "" {
t.Errorf("ID = %q, want \"\"", got.ID)
}
if got.DocumentID != "" {
t.Errorf("DocumentID = %q, want \"\"", got.DocumentID)
}
}
// TestTranslateChunk_MissingAllScores: a chunk with no score
// fields at all → score 0 (don't panic).
func TestTranslateChunk_MissingAllScores(t *testing.T) {
raw := map[string]any{"chunk_id": "ck-1"}
got := translateChunk(raw)
if got.Score != 0 {
t.Errorf("Score = %v, want 0", got.Score)
}
}
// TestNewNLPRetrievalAdapter_NilService: nil constructor inputs
// must produce an adapter whose Search returns the missing-service
// error, not a panic.
func TestNewNLPRetrievalAdapter_NilService(t *testing.T) {
a := NewNLPRetrievalAdapter(nil, nil, nil)
_, err := a.Search(context.TODO(), nil, RetrievalRequest{Query: "hi"})
if err == nil {
t.Fatal("expected error from nil-service adapter")
}
if !errors.Is(err, ErrRetrievalServiceMissing) {
t.Errorf("err = %v, want ErrRetrievalServiceMissing", err)
}
}
func TestNLPRequestFromRetrieval_ThreadsSearchControls(t *testing.T) {
keywordWeight := 0.25
similarityThreshold := 0.42
embeddingModel := &modelModule.EmbeddingModel{}
got := nlpRequestFromRetrieval(RetrievalRequest{
Query: "hi",
DatasetIDs: []string{"kb-1"},
TopN: 3,
TopK: 99,
KeywordsSimilarityWeight: &keywordWeight,
SimilarityThreshold: &similarityThreshold,
}, []string{"tenant-a"}, 3, embeddingModel)
if got.Question != "hi" {
t.Fatalf("Question=%q want hi", got.Question)
}
if len(got.TenantIDs) != 1 || got.TenantIDs[0] != "tenant-a" {
t.Fatalf("TenantIDs=%v want [tenant-a]", got.TenantIDs)
}
if len(got.KbIDs) != 1 || got.KbIDs[0] != "kb-1" {
t.Fatalf("KbIDs=%v want [kb-1]", got.KbIDs)
}
if got.Page != 1 || got.PageSize != 3 {
t.Fatalf("Page/PageSize=%d/%d want 1/3", got.Page, got.PageSize)
}
if got.Top == nil || *got.Top != 99 {
t.Fatalf("Top=%v want 99", got.Top)
}
if got.SimilarityThreshold == nil || *got.SimilarityThreshold != 0.42 {
t.Fatalf("SimilarityThreshold=%v want 0.42", got.SimilarityThreshold)
}
if got.VectorSimilarityWeight == nil || !floatEqual(*got.VectorSimilarityWeight, 0.75) {
t.Fatalf("VectorSimilarityWeight=%v want 0.75", got.VectorSimilarityWeight)
}
if got.EmbeddingModel != embeddingModel {
t.Fatal("EmbeddingModel was not passed to nlp retrieval request")
}
}
func TestNLPRequestFromRetrieval_FallsBackToTopNHeadroom(t *testing.T) {
got := nlpRequestFromRetrieval(RetrievalRequest{
Query: "hi",
DatasetIDs: []string{"kb-1"},
TopN: 3,
}, []string{"tenant-a"}, 3, &modelModule.EmbeddingModel{})
if got.Top == nil || *got.Top != 12 {
t.Fatalf("Top=%v want 12", got.Top)
}
if got.VectorSimilarityWeight != nil {
t.Fatalf("VectorSimilarityWeight=%v want nil", got.VectorSimilarityWeight)
}
}
func TestNLPRequestFromRetrieval_PreservesExplicitZeroSimilarityThreshold(t *testing.T) {
similarityThreshold := 0.0
got := nlpRequestFromRetrieval(RetrievalRequest{
Query: "hi",
DatasetIDs: []string{"kb-1"},
SimilarityThreshold: &similarityThreshold,
}, []string{"tenant-a"}, 3, &modelModule.EmbeddingModel{})
if got.SimilarityThreshold == nil || *got.SimilarityThreshold != 0 {
t.Fatalf("SimilarityThreshold = %v; want explicit zero", got.SimilarityThreshold)
}
}
func TestNLPRetrievalAdapter_ResolveDatasetsRequiresEveryDataset(t *testing.T) {
a := &NLPRetrievalAdapter{
kbDAO: fakeKnowledgebaseLookup{
kbs: []*entity.Knowledgebase{
{ID: "kb-1", TenantID: "tenant-a"},
{ID: "kb-2", TenantID: "tenant-a"},
},
},
}
_, err := a.resolveDatasets(nil, nil, RetrievalRequest{
TenantID: "tenant-a",
DatasetIDs: []string{"kb-1", "kb-2", "kb-missing"},
})
if err == nil {
t.Fatal("resolveDatasets: expected missing dataset error")
}
}
func TestNLPRetrievalAdapter_SearchRejectsDatasetsFromMultipleTenants(t *testing.T) {
a := &NLPRetrievalAdapter{
svc: &nlp.RetrievalService{},
kbDAO: fakeKnowledgebaseLookup{
kbs: []*entity.Knowledgebase{
{ID: "kb-1", TenantID: "tenant-a", EmbdID: "embedding@provider"},
{ID: "kb-2", TenantID: "tenant-b", EmbdID: "embedding@provider"},
},
},
}
_, err := a.Search(t.Context(), nil, RetrievalRequest{
Query: "hi",
DatasetIDs: []string{"kb-1", "kb-2"},
})
if err == nil || err.Error() != "retrieval: datasets span multiple tenants" {
t.Fatalf("Search error = %v, want multiple-tenant rejection", err)
}
}
func TestNLPRetrievalAdapter_ResolveDatasetsFromDatasetIDs(t *testing.T) {
a := &NLPRetrievalAdapter{
kbDAO: fakeKnowledgebaseLookup{
kbs: []*entity.Knowledgebase{
{ID: "kb-1", TenantID: "tenant-a"},
{ID: "kb-2", TenantID: "tenant-b"},
{ID: "kb-3", TenantID: "tenant-a"},
},
},
}
got, err := a.resolveDatasets(nil, nil, RetrievalRequest{
DatasetIDs: []string{"kb-1", "kb-2", "kb-3", " "},
})
if err != nil {
t.Fatalf("resolveDatasets: %v", err)
}
if len(got.tenantIDs) != 2 || got.tenantIDs[0] != "tenant-a" || got.tenantIDs[1] != "tenant-b" {
t.Fatalf("tenantIDs=%v want [tenant-a tenant-b]", got.tenantIDs)
}
}
func TestNLPRetrievalAdapter_ResolveDatasetsUsesDatasetTenants(t *testing.T) {
a := &NLPRetrievalAdapter{
kbDAO: fakeKnowledgebaseLookup{
kbs: []*entity.Knowledgebase{
{ID: "kb-1", TenantID: "tenant-b"},
},
},
}
got, err := a.resolveDatasets(nil, nil, RetrievalRequest{
TenantID: "tenant-a",
DatasetIDs: []string{"kb-1"},
})
if err != nil {
t.Fatalf("resolveDatasets: %v", err)
}
if len(got.tenantIDs) != 1 || got.tenantIDs[0] != "tenant-b" {
t.Fatalf("tenantIDs=%v want [tenant-b]", got.tenantIDs)
}
}
func TestNLPRetrievalAdapter_ResolveEmbeddingModelPriority(t *testing.T) {
tenantEmbeddingID := "tenant-embedding-1"
tests := []struct {
name string
kb *entity.Knowledgebase
wantCall string
}{
{
name: "tenant embedding id",
kb: &entity.Knowledgebase{
ID: "kb-1",
TenantID: "tenant-1",
EmbdID: "embedding@provider",
TenantEmbdID: &tenantEmbeddingID,
},
wantCall: "id:tenant-embedding-1",
},
{
name: "knowledge base embedding reference",
kb: &entity.Knowledgebase{
ID: "kb-1",
TenantID: "tenant-1",
EmbdID: "embedding@provider",
},
wantCall: "resolve:embedding@provider",
},
{
name: "tenant default only when knowledge base has no model",
kb: &entity.Knowledgebase{
ID: "kb-1",
TenantID: "tenant-1",
},
wantCall: "default",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
resolver := &fakeModelResolver{modelName: "resolved-model"}
adapter := &NLPRetrievalAdapter{modelResolver: resolver}
model, err := adapter.resolveEmbeddingModel(t.Context(), test.kb)
if err != nil {
t.Fatalf("resolveEmbeddingModel: %v", err)
}
if resolver.call != test.wantCall {
t.Fatalf("resolver call = %q, want %q", resolver.call, test.wantCall)
}
if model == nil || model.ModelName == nil || *model.ModelName != "resolved-model" {
t.Fatalf("resolved model = %#v", model)
}
})
}
}
func TestValidateEmbeddingModelsRejectsDifferentTenantModels(t *testing.T) {
firstID := "tenant-embedding-1"
secondID := "tenant-embedding-2"
err := validateEmbeddingModels([]*entity.Knowledgebase{
{ID: "kb-1", TenantID: "tenant-1", TenantEmbdID: &firstID},
{ID: "kb-2", TenantID: "tenant-1", TenantEmbdID: &secondID},
})
if err == nil {
t.Fatal("expected different tenant embedding models to be rejected")
}
}
type fakeKnowledgebaseLookup struct {
kbs []*entity.Knowledgebase
err error
}
type fakeModelResolver struct {
call string
modelName string
err error
}
func (f *fakeModelResolver) GetModelConfigByID(
_ context.Context,
_ string,
_ entity.ModelType,
modelID string,
) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) {
f.call = "id:" + modelID
return nil, f.modelName, &modelModule.APIConfig{}, 512, f.err
}
func (f *fakeModelResolver) ResolveModelConfig(
_ context.Context,
_ string,
_ entity.ModelType,
modelRef string,
) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) {
f.call = "resolve:" + modelRef
return nil, f.modelName, &modelModule.APIConfig{}, 512, f.err
}
func (f *fakeModelResolver) GetTenantDefaultModelByType(
_ context.Context,
_ string,
_ entity.ModelType,
) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) {
f.call = "default"
return nil, f.modelName, &modelModule.APIConfig{}, 512, f.err
}
func (f fakeKnowledgebaseLookup) GetByIDs(ctx context.Context, db *gorm.DB, ids []string) ([]*entity.Knowledgebase, error) {
requested := make(map[string]struct{}, len(ids))
for _, id := range ids {
requested[id] = struct{}{}
}
kbs := make([]*entity.Knowledgebase, 0, len(f.kbs))
for _, kb := range f.kbs {
if kb == nil {
continue
}
if _, ok := requested[kb.ID]; ok {
kbs = append(kbs, kb)
}
}
return kbs, f.err
}
func (f fakeKnowledgebaseLookup) GetByName(_ context.Context, _ *gorm.DB, name, tenantID string) (*entity.Knowledgebase, error) {
for _, kb := range f.kbs {
if kb != nil && kb.Name == name && kb.TenantID == tenantID {
return kb, nil
}
}
return nil, gorm.ErrRecordNotFound
}