Files

78 lines
1.7 KiB
Go

package index
import (
"context"
"eino-test/internal/embedding"
"eino-test/internal/store"
"github.com/cloudwego/eino/components/retriever"
"github.com/cloudwego/eino/schema"
)
type HybridRetriever struct {
vectorIdx *InMemoryVectorIndex
embedder *embedding.Embedder
noteStore store.NoteStore
topK int
}
func NewHybridRetriever(vectorIdx *InMemoryVectorIndex, embedder *embedding.Embedder, noteStore store.NoteStore, topK int) *HybridRetriever {
return &HybridRetriever{
vectorIdx: vectorIdx,
embedder: embedder,
noteStore: noteStore,
topK: topK,
}
}
func (r *HybridRetriever) Retrieve(ctx context.Context, query string, opts ...retriever.Option) ([]*schema.Document, error) {
options := &retriever.Options{
TopK: &r.topK,
}
options = retriever.GetCommonOptions(options, opts...)
topK := r.topK
if options.TopK != nil {
topK = *options.TopK
}
queryVecs, err := r.embedder.Embed(ctx, []string{query})
if err != nil {
return nil, err
}
hits := r.vectorIdx.Search(queryVecs[0], topK)
var docs []*schema.Document
for _, hit := range hits {
if options.ScoreThreshold != nil && hit.Score < *options.ScoreThreshold {
continue
}
note, err := r.noteStore.GetByID(ctx, noteIDFromChunkID(hit.ID))
if err != nil {
continue
}
docs = append(docs, &schema.Document{
ID: hit.ID,
Content: note.Content,
MetaData: map[string]any{
"note_id": note.ID,
"title": note.Title,
"tags": note.Tags,
"score": hit.Score,
},
})
}
return docs, nil
}
func noteIDFromChunkID(chunkID string) string {
for i := len(chunkID) - 1; i >= 0; i-- {
if chunkID[i] == '_' {
return chunkID[:i]
}
}
return chunkID
}