Merge branch 'pr/747' into main-pull-requests

This commit is contained in:
Dmitry Ng
2025-05-06 18:49:41 +03:00
2 changed files with 85 additions and 0 deletions
+45
View File
@@ -0,0 +1,45 @@
package retrievers
import (
"context"
"github.com/tmc/langchaingo/callbacks"
"github.com/tmc/langchaingo/schema"
)
var _ schema.Retriever = &MergerRetriever{}
// MergerRetriever is a retriever that merges the results of multiple retrievers.
type MergerRetriever struct {
Retrievers []schema.Retriever
CallbacksHandler callbacks.Handler
}
// NewMergerRetriever creates a new MergerRetriever.
func NewMergerRetriever(
retrievers []schema.Retriever,
) MergerRetriever {
return MergerRetriever{
Retrievers: retrievers,
CallbacksHandler: nil,
}
}
// GetRelevantDocuments returns documents from the MergerRetriever's all retrievers.
func (m *MergerRetriever) GetRelevantDocuments(ctx context.Context, query string) ([]schema.Document, error) {
if m.CallbacksHandler != nil {
m.CallbacksHandler.HandleRetrieverStart(ctx, query)
}
docs := make([]schema.Document, 0)
for _, r := range m.Retrievers {
doc, err := r.GetRelevantDocuments(ctx, query)
if err != nil {
return nil, err
}
docs = append(docs, doc...)
}
if m.CallbacksHandler != nil {
m.CallbacksHandler.HandleRetrieverEnd(ctx, query, docs)
}
return docs, nil
}
+40
View File
@@ -0,0 +1,40 @@
package retrievers
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/tmc/langchaingo/schema"
)
var _ schema.Retriever = &Fakeretriever{}
type Fakeretriever struct {
Docs []schema.Document
}
func (f *Fakeretriever) GetRelevantDocuments(_ context.Context, _ string) ([]schema.Document, error) {
return f.Docs, nil
}
func TestMergerRetriever(t *testing.T) { //nolint:funlen
t.Parallel()
ctx := context.Background()
content1 := schema.Document{PageContent: "fake doc 1"}
content2 := schema.Document{PageContent: "fake doc 2"}
content3 := schema.Document{PageContent: "fake doc 3"}
content4 := schema.Document{PageContent: "fake doc 4"}
retriever1 := Fakeretriever{Docs: []schema.Document{content1, content2}}
retriever2 := Fakeretriever{Docs: []schema.Document{content3, content4}}
merger := NewMergerRetriever([]schema.Retriever{&retriever1, &retriever2})
documents, err := merger.GetRelevantDocuments(ctx, "fake query")
require.NoError(t, err)
require.Len(t, documents, 4)
assert.Equal(t, documents[0], content1)
assert.Equal(t, documents[1], content2)
assert.Equal(t, documents[2], content3)
assert.Equal(t, documents[3], content4)
}