mirror of
https://github.com/vxcontrol/langchaingo.git
synced 2026-07-21 00:45:22 -04:00
Merge branch 'pr/747' into main-pull-requests
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user