Enhance Collection Retrieval and Existence Check Logic

- Updated the `GetCollection` method to first read from the database for existence and permissions, improving data integrity.
- Merged metadata from GraphRag into the result, ensuring backward compatibility and enhanced data representation.
- Refactored the `CollectionExists` method to check both the database and GraphRag for consistency, logging any mismatches for debugging purposes.
- Introduced new types and structures for search operations, including `SearchMode`, `Query`, and `SearchResult`, to support advanced search functionalities.
This commit is contained in:
Max 2025-12-20 11:25:02 +08:00
parent fa471bc4c4
commit 53db8be522
6 changed files with 1521 additions and 29 deletions

View file

@ -274,61 +274,79 @@ func (instance *KBInstance) RemoveCollection(ctx context.Context, collectionID s
}
// GetCollection retrieves a collection by ID
// Reads from database first, then merges with GraphRag metadata
func (instance *KBInstance) GetCollection(ctx context.Context, collectionID string) (map[string]interface{}, error) {
if collectionID == "" {
return nil, fmt.Errorf("collection ID is required")
}
collection, err := instance.GraphRag.GetCollection(ctx, collectionID)
// Read from database (source of truth for existence and permissions)
dbRecord, err := instance.Config.FindCollection(collectionID, model.QueryParam{})
if err != nil {
// Check if it's a "not found" error
if err.Error() == fmt.Sprintf("collection with ID '%s' not found", collectionID) {
return nil, fmt.Errorf("collection not found")
}
return nil, fmt.Errorf("failed to get collection: %w", err)
return nil, fmt.Errorf("collection not found")
}
// Convert CollectionInfo to map[string]interface{}
// Use a hybrid structure: flatten metadata to top level AND include metadata object
// This ensures backward compatibility with both access patterns:
// - collection.id / collection.collection_id (for ID)
// - collection.metadata.name (for nested access)
// Convert database record to result map (flatten to top level)
result := make(map[string]interface{})
result["id"] = collection.ID // Primary ID field for frontend
result["collection_id"] = collection.ID // Alias for backward compatibility
// Flatten metadata fields to top level for backward compatibility
if collection.Metadata != nil {
for k, v := range collection.Metadata {
result[k] = v
}
// Also include the metadata object itself
result["metadata"] = collection.Metadata
for k, v := range dbRecord {
result[k] = v
}
if collection.Config != nil {
result["config"] = collection.Config
// Set standard ID fields
result["id"] = collectionID
result["collection_id"] = collectionID
// Read from GraphRag and merge (for config and metadata object)
graphRagCollection, err := instance.GraphRag.GetCollection(ctx, collectionID)
if err == nil && graphRagCollection != nil {
// Set GraphRag config (vector store configuration)
if graphRagCollection.Config != nil {
result["config"] = graphRagCollection.Config
}
// Set GraphRag metadata as nested object (for backward compatibility)
// This allows access via collection["metadata"]["field"]
if graphRagCollection.Metadata != nil {
result["metadata"] = graphRagCollection.Metadata
// Also flatten GraphRag metadata fields to top level
// Only add fields that don't exist in database record
for k, v := range graphRagCollection.Metadata {
if _, exists := result[k]; !exists {
result[k] = v
}
}
}
}
return result, nil
}
// CollectionExists checks if a collection exists by ID
// Checks both database and GraphRag for consistency
func (instance *KBInstance) CollectionExists(ctx context.Context, collectionID string) (*CollectionExistsResult, error) {
if collectionID == "" {
return nil, fmt.Errorf("collection ID is required")
}
exists, err := instance.GraphRag.CollectionExists(ctx, collectionID)
if err != nil {
return nil, fmt.Errorf("failed to check collection existence: %w", err)
// Check database (source of truth for existence)
_, dbErr := instance.Config.FindCollection(collectionID, model.QueryParam{})
dbExists := dbErr == nil
// Check GraphRag for consistency
graphRagExists, _ := instance.GraphRag.CollectionExists(ctx, collectionID)
// Collection exists if it exists in database
// Log warning if there's inconsistency (for debugging)
if dbExists != graphRagExists {
log.Warn("Collection %s existence mismatch: database=%v, graphrag=%v", collectionID, dbExists, graphRagExists)
}
return &CollectionExistsResult{
CollectionID: collectionID,
Exists: exists,
Exists: dbExists,
}, nil
}

View file

@ -32,8 +32,8 @@ type API interface {
AddTextAsync(ctx context.Context, params *AddTextParams) (*AddDocumentAsyncResult, error)
AddURLAsync(ctx context.Context, params *AddURLParams) (*AddDocumentAsyncResult, error)
// Segment operations (future)
// ...
// Search operations
Search(ctx context.Context, queries []Query) (*SearchResult, error)
}
// KBInstance holds the KB instance dependencies required by the API

551
kb/api/search.go Normal file
View file

@ -0,0 +1,551 @@
package api
import (
"context"
"fmt"
"sort"
"sync"
graphragtypes "github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/yao/kb/providers/factory"
)
// Default search parameters
const (
DefaultSearchK = 10
DefaultMaxDepth = 2
DefaultMinScore = 0.0
MaxSearchK = 100
DefaultSearchPageSize = 20
)
// Search performs batch search operations on the knowledge base
// Queries can span multiple collections; implementation groups by CollectionID
// Mode, providers (embedding/extraction/reranker) are read from each collection's config
// All results are merged and deduplicated
func (kb *KBInstance) Search(ctx context.Context, queries []Query) (*SearchResult, error) {
if len(queries) == 0 {
return &SearchResult{
Segments: []graphragtypes.Segment{},
Total: 0,
}, nil
}
// 1. Validate queries
if err := kb.validateQueries(queries); err != nil {
return nil, err
}
// 2. Group queries by CollectionID
groupedQueries := kb.groupQueriesByCollection(queries)
// 3. Process each collection group in parallel
var (
allSegments []graphragtypes.Segment
allGraph *GraphData
mu sync.Mutex
wg sync.WaitGroup
errChan = make(chan error, len(groupedQueries))
)
for collectionID, collQueries := range groupedQueries {
wg.Add(1)
go func(collID string, qs []Query) {
defer wg.Done()
segments, graph, err := kb.searchCollection(ctx, collID, qs)
if err != nil {
errChan <- fmt.Errorf("search in collection %s failed: %w", collID, err)
return
}
mu.Lock()
allSegments = append(allSegments, segments...)
if graph != nil {
allGraph = mergeGraphData(allGraph, graph)
}
mu.Unlock()
}(collectionID, collQueries)
}
wg.Wait()
close(errChan)
// Collect errors
var errors []error
for err := range errChan {
errors = append(errors, err)
}
if len(errors) > 0 {
log.Warn("Search completed with errors: %v", errors)
}
// 4. Merge and deduplicate results
mergedSegments := kb.deduplicateSegments(allSegments)
// 5. Sort by score (descending)
sort.Slice(mergedSegments, func(i, j int) bool {
return mergedSegments[i].Score > mergedSegments[j].Score
})
// 6. Apply pagination from first query (if specified)
result := kb.applyPagination(mergedSegments, queries[0])
result.Graph = allGraph
return result, nil
}
// ========== Validation ==========
// validateQueries validates all queries
func (kb *KBInstance) validateQueries(queries []Query) error {
for i, q := range queries {
if q.CollectionID == "" {
return fmt.Errorf("query %d: collection_id is required", i)
}
if q.Input == "" && len(q.Messages) == 0 {
return fmt.Errorf("query %d: either input or messages is required", i)
}
}
return nil
}
// ========== Query Grouping ==========
// groupQueriesByCollection groups queries by their CollectionID
func (kb *KBInstance) groupQueriesByCollection(queries []Query) map[string][]Query {
grouped := make(map[string][]Query)
for _, q := range queries {
grouped[q.CollectionID] = append(grouped[q.CollectionID], q)
}
return grouped
}
// ========== Collection Search ==========
// searchCollection processes all queries for a single collection
func (kb *KBInstance) searchCollection(ctx context.Context, collectionID string, queries []Query) ([]graphragtypes.Segment, *GraphData, error) {
// Get collection config
collection, err := kb.GetCollection(ctx, collectionID)
if err != nil {
return nil, nil, fmt.Errorf("failed to get collection: %w", err)
}
// Get embedding provider from collection config
embeddingProviderID, _ := collection["embedding_provider_id"].(string)
embeddingOptionID, _ := collection["embedding_option_id"].(string)
if embeddingProviderID == "" || embeddingOptionID == "" {
return nil, nil, fmt.Errorf("collection %s missing embedding configuration", collectionID)
}
// Create embedding function
embedding, err := kb.createEmbedding(embeddingProviderID, embeddingOptionID, "en")
if err != nil {
return nil, nil, fmt.Errorf("failed to create embedding: %w", err)
}
var (
allSegments []graphragtypes.Segment
allGraph *GraphData
mu sync.Mutex
wg sync.WaitGroup
errChan = make(chan error, len(queries))
)
// Process queries in parallel
for _, query := range queries {
wg.Add(1)
go func(q Query) {
defer wg.Done()
segments, graph, err := kb.executeQuery(ctx, collectionID, q, embedding, collection)
if err != nil {
errChan <- err
return
}
mu.Lock()
allSegments = append(allSegments, segments...)
if graph != nil {
allGraph = mergeGraphData(allGraph, graph)
}
mu.Unlock()
}(query)
}
wg.Wait()
close(errChan)
// Collect errors
var errors []error
for err := range errChan {
errors = append(errors, err)
}
if len(errors) > 0 {
return allSegments, allGraph, errors[0]
}
return allSegments, allGraph, nil
}
// executeQuery executes a single query based on its mode
func (kb *KBInstance) executeQuery(ctx context.Context, collectionID string, query Query, embedding graphragtypes.Embedding, collection map[string]interface{}) ([]graphragtypes.Segment, *GraphData, error) {
// Determine search mode
mode := query.Mode
if mode == "" {
// Default to expand mode
mode = SearchModeExpand
}
// Get query text
queryText := kb.getQueryText(query)
if queryText == "" {
return nil, nil, fmt.Errorf("no query text found")
}
// Execute based on mode
switch mode {
case SearchModeVector:
return kb.searchVector(ctx, collectionID, queryText, query, embedding)
case SearchModeGraph:
return kb.searchGraph(ctx, collectionID, queryText, query, collection)
case SearchModeExpand:
return kb.searchExpand(ctx, collectionID, queryText, query, embedding, collection)
default:
return nil, nil, fmt.Errorf("unknown search mode: %s", mode)
}
}
// getQueryText extracts query text from Input or Messages
func (kb *KBInstance) getQueryText(query Query) string {
// Input takes precedence
if query.Input != "" {
return query.Input
}
// Extract from last user message
for i := len(query.Messages) - 1; i >= 0; i-- {
if query.Messages[i].Role == "user" {
return query.Messages[i].Content
}
}
return ""
}
// ========== Vector Search ==========
// searchVector performs pure vector similarity search
func (kb *KBInstance) searchVector(ctx context.Context, collectionID string, queryText string, query Query, embedding graphragtypes.Embedding) ([]graphragtypes.Segment, *GraphData, error) {
// Build search options
k := query.PageSize
if k <= 0 {
k = DefaultSearchK
}
if k > MaxSearchK {
k = MaxSearchK
}
options := &graphragtypes.VectorSearchOptions{
CollectionID: collectionID,
DocumentID: query.DocumentID,
Query: queryText,
K: k,
MinScore: query.MinScore,
Embedding: embedding,
}
// Add metadata filter
if len(query.Metadata) > 0 {
options.Filter = query.Metadata
}
// Execute search
result, err := kb.GraphRag.SearchVector(ctx, options)
if err != nil {
return nil, nil, fmt.Errorf("vector search failed: %w", err)
}
return result.Segments, nil, nil
}
// ========== Graph Search ==========
// searchGraph performs pure graph traversal search
func (kb *KBInstance) searchGraph(ctx context.Context, collectionID string, queryText string, query Query, collection map[string]interface{}) ([]graphragtypes.Segment, *GraphData, error) {
// Get extraction provider for entity extraction
extraction, err := kb.createExtraction(collection)
if err != nil {
return nil, nil, fmt.Errorf("failed to create extraction: %w", err)
}
// Build graph search options
maxDepth := query.MaxDepth
if maxDepth <= 0 {
maxDepth = DefaultMaxDepth
}
options := &graphragtypes.GraphSearchOptions{
CollectionID: collectionID,
DocumentID: query.DocumentID,
Query: queryText,
MaxDepth: maxDepth,
Extraction: extraction,
}
// Execute search
result, err := kb.GraphRag.SearchGraph(ctx, options)
if err != nil {
return nil, nil, fmt.Errorf("graph search failed: %w", err)
}
// Convert to GraphData
graph := &GraphData{
Nodes: result.Nodes,
Relationships: result.Relationships,
}
return result.Segments, graph, nil
}
// ========== Expand Search (Graph + Vector) ==========
// searchExpand performs graph-based entity expansion + vector search
// This mode uses graph to find related entities, then enhances vector search
func (kb *KBInstance) searchExpand(ctx context.Context, collectionID string, queryText string, query Query, embedding graphragtypes.Embedding, collection map[string]interface{}) ([]graphragtypes.Segment, *GraphData, error) {
// Step 1: Extract entities from query using graph search
extraction, err := kb.createExtraction(collection)
if err != nil {
// Fall back to pure vector search if extraction is not available
log.Warn("Extraction not available, falling back to vector search: %v", err)
return kb.searchVector(ctx, collectionID, queryText, query, embedding)
}
maxDepth := query.MaxDepth
if maxDepth <= 0 {
maxDepth = DefaultMaxDepth
}
graphOptions := &graphragtypes.GraphSearchOptions{
CollectionID: collectionID,
DocumentID: query.DocumentID,
Query: queryText,
MaxDepth: maxDepth,
Extraction: extraction,
}
// Execute graph search to find related entities
graphResult, graphErr := kb.GraphRag.SearchGraph(ctx, graphOptions)
// Step 2: Perform vector search
k := query.PageSize
if k <= 0 {
k = DefaultSearchK
}
if k > MaxSearchK {
k = MaxSearchK
}
vectorOptions := &graphragtypes.VectorSearchOptions{
CollectionID: collectionID,
DocumentID: query.DocumentID,
Query: queryText,
K: k,
MinScore: query.MinScore,
Embedding: embedding,
}
if len(query.Metadata) > 0 {
vectorOptions.Filter = query.Metadata
}
vectorResult, err := kb.GraphRag.SearchVector(ctx, vectorOptions)
if err != nil {
return nil, nil, fmt.Errorf("vector search failed: %w", err)
}
// Step 3: Merge results
segments := vectorResult.Segments
var graph *GraphData
if graphErr == nil && graphResult != nil {
// Add graph segments (deduplicated later)
segments = append(segments, graphResult.Segments...)
// Include graph data
graph = &GraphData{
Nodes: graphResult.Nodes,
Relationships: graphResult.Relationships,
}
}
return segments, graph, nil
}
// ========== Helper Functions ==========
// createEmbedding creates an embedding function from provider config
func (kb *KBInstance) createEmbedding(providerID, optionID, locale string) (graphragtypes.Embedding, error) {
if locale == "" {
locale = "en"
}
// Get provider option
option, err := kb.getProviderOption("embedding", providerID, optionID, locale)
if err != nil {
return nil, fmt.Errorf("failed to get embedding option: %w", err)
}
// Create embedding provider
return factory.MakeEmbedding(providerID, option)
}
// createExtraction creates an extraction function from collection config
func (kb *KBInstance) createExtraction(collection map[string]interface{}) (graphragtypes.Extraction, error) {
// Try to get extraction provider from collection metadata
metadata, _ := collection["metadata"].(map[string]interface{})
if metadata == nil {
metadata = collection
}
extractionProviderID, _ := metadata["__extraction_provider"].(string)
extractionOptionID, _ := metadata["__extraction_option"].(string)
// Fall back to default extraction provider
if extractionProviderID == "" {
extractionProviderID = "__yao.openai"
extractionOptionID = "gpt-4o-mini"
}
// Get provider option
option, err := kb.getProviderOption("extraction", extractionProviderID, extractionOptionID, "en")
if err != nil {
return nil, fmt.Errorf("failed to get extraction option: %w", err)
}
// Create extraction provider
return factory.MakeExtraction(extractionProviderID, option)
}
// deduplicateSegments removes duplicate segments by ID, keeping highest score
func (kb *KBInstance) deduplicateSegments(segments []graphragtypes.Segment) []graphragtypes.Segment {
seen := make(map[string]int) // ID -> index in result
result := make([]graphragtypes.Segment, 0, len(segments))
for _, seg := range segments {
if idx, exists := seen[seg.ID]; exists {
// Keep the one with higher score
if seg.Score > result[idx].Score {
result[idx] = seg
}
} else {
seen[seg.ID] = len(result)
result = append(result, seg)
}
}
return result
}
// mergeGraphData merges two GraphData objects
func mergeGraphData(a, b *GraphData) *GraphData {
if a == nil {
return b
}
if b == nil {
return a
}
// Merge nodes (deduplicate by ID)
nodeMap := make(map[string]graphragtypes.GraphNode)
for _, n := range a.Nodes {
nodeMap[n.ID] = n
}
for _, n := range b.Nodes {
nodeMap[n.ID] = n
}
nodes := make([]graphragtypes.GraphNode, 0, len(nodeMap))
for _, n := range nodeMap {
nodes = append(nodes, n)
}
// Merge relationships (deduplicate by ID)
relMap := make(map[string]graphragtypes.GraphRelationship)
for _, r := range a.Relationships {
relMap[r.ID] = r
}
for _, r := range b.Relationships {
relMap[r.ID] = r
}
relationships := make([]graphragtypes.GraphRelationship, 0, len(relMap))
for _, r := range relMap {
relationships = append(relationships, r)
}
return &GraphData{
Nodes: nodes,
Relationships: relationships,
}
}
// applyPagination applies pagination to segments
func (kb *KBInstance) applyPagination(segments []graphragtypes.Segment, query Query) *SearchResult {
total := len(segments)
// If no pagination requested, return all
if query.Page <= 0 && query.PageSize <= 0 {
return &SearchResult{
Segments: segments,
Total: total,
}
}
page := query.Page
if page <= 0 {
page = 1
}
pageSize := query.PageSize
if pageSize <= 0 {
pageSize = DefaultSearchPageSize
}
// Calculate pagination
totalPages := (total + pageSize - 1) / pageSize
start := (page - 1) * pageSize
end := start + pageSize
if start >= total {
return &SearchResult{
Segments: []graphragtypes.Segment{},
Total: total,
Page: page,
PageSize: pageSize,
TotalPages: totalPages,
}
}
if end > total {
end = total
}
result := &SearchResult{
Segments: segments[start:end],
Total: total,
Page: page,
PageSize: pageSize,
TotalPages: totalPages,
}
// Set next/prev page
if page < totalPages {
result.Next = page + 1
}
if page > 1 {
result.Prev = page - 1
}
return result
}

400
kb/api/search_setup_test.go Normal file
View file

@ -0,0 +1,400 @@
package api_test
import (
"context"
"fmt"
"testing"
"time"
"github.com/stretchr/testify/assert"
graphragtypes "github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/yao/kb"
"github.com/yaoapp/yao/kb/api"
)
// Note: TestMain is defined in collection_test.go
// ========== Fixed Test Collection IDs ==========
// Use fixed IDs so we can reuse them across test runs during development
const (
// SearchTestScienceCollection is the fixed ID for science test collection
SearchTestScienceCollection = "search_test_science"
// SearchTestTechCollection is the fixed ID for tech test collection
SearchTestTechCollection = "search_test_tech"
)
// ========== Setup Test - Run Once ==========
// TestSearchSetup creates test collections and documents for search testing.
// Run this once before running search tests:
//
// go test -v -run "TestSearchSetup" ./kb/api/...
//
// Then run search tests multiple times without waiting for data setup:
//
// go test -v -run "TestSearchQuery" ./kb/api/...
func TestSearchSetup(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
// Check if collections already exist and are complete
// We check both GraphRag (vector store) and document count
scienceComplete := false
techComplete := false
// Check Science collection
scienceCollection, scienceErr := kb.API.GetCollection(ctx, SearchTestScienceCollection)
if scienceErr == nil && scienceCollection != nil {
scienceDocs, _ := kb.API.ListDocuments(ctx, &api.ListDocumentsFilter{
Page: 1,
PageSize: 20,
CollectionID: SearchTestScienceCollection,
})
if scienceDocs != nil && len(scienceDocs.Data) >= 5 {
scienceComplete = true
t.Logf("✓ Science collection exists: %s (%d docs)", SearchTestScienceCollection, len(scienceDocs.Data))
}
}
// Check Tech collection
techCollection, techErr := kb.API.GetCollection(ctx, SearchTestTechCollection)
if techErr == nil && techCollection != nil {
techDocs, _ := kb.API.ListDocuments(ctx, &api.ListDocumentsFilter{
Page: 1,
PageSize: 20,
CollectionID: SearchTestTechCollection,
})
if techDocs != nil && len(techDocs.Data) >= 5 {
techComplete = true
t.Logf("✓ Tech collection exists: %s (%d docs)", SearchTestTechCollection, len(techDocs.Data))
}
}
// If both collections are complete, skip setup
if scienceComplete && techComplete {
t.Log("✓ All test collections already exist with sufficient documents")
t.Log(" Skipping setup. Run TestSearchCleanup first to recreate.")
return
}
// Clean up any existing collections (handles both complete and incomplete states)
// RemoveCollection cleans both database and GraphRag (including orphaned vector collections)
t.Log("Cleaning up existing collections...")
if result, err := kb.API.RemoveCollection(ctx, SearchTestScienceCollection); err == nil && result.Removed {
t.Logf(" Removed: %s", SearchTestScienceCollection)
}
if result, err := kb.API.RemoveCollection(ctx, SearchTestTechCollection); err == nil && result.Removed {
t.Logf(" Removed: %s", SearchTestTechCollection)
}
time.Sleep(1 * time.Second) // Wait for cleanup
// Create Science Collection
t.Log("Creating Science collection...")
scienceParams := &api.CreateCollectionParams{
ID: SearchTestScienceCollection,
Metadata: map[string]interface{}{
"name": "Science Knowledge Base",
"description": "Scientists and their discoveries for search testing",
},
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Locale: "en",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
_, err := kb.API.CreateCollection(ctx, scienceParams)
if err != nil {
t.Fatalf("Failed to create science collection: %v", err)
}
t.Logf("✓ Created collection: %s", SearchTestScienceCollection)
// Create Tech Collection
t.Log("Creating Tech collection...")
techParams := &api.CreateCollectionParams{
ID: SearchTestTechCollection,
Metadata: map[string]interface{}{
"name": "Tech Knowledge Base",
"description": "Technology companies and products for search testing",
},
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Locale: "en",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
_, err = kb.API.CreateCollection(ctx, techParams)
if err != nil {
t.Fatalf("Failed to create tech collection: %v", err)
}
t.Logf("✓ Created collection: %s", SearchTestTechCollection)
// Add Science Documents
// Entity relationships: Einstein -> Relativity -> Physics -> Nobel Prize
scienceDocs := []struct {
title string
content string
}{
{
title: "Albert Einstein Biography",
content: `Albert Einstein was a theoretical physicist born in Germany in 1879.
He developed the theory of relativity, one of the two pillars of modern physics.
Einstein received the Nobel Prize in Physics in 1921 for his discovery of the photoelectric effect.
He later emigrated to the United States and worked at Princeton University until his death in 1955.`,
},
{
title: "Theory of Relativity",
content: `The theory of relativity was developed by Albert Einstein in the early 20th century.
It consists of special relativity (1905) and general relativity (1915).
Special relativity introduced E=mc², showing the relationship between energy and mass.
General relativity describes gravity as the curvature of spacetime caused by mass and energy.`,
},
{
title: "Marie Curie Biography",
content: `Marie Curie was a Polish-French physicist and chemist who conducted pioneering research on radioactivity.
She was the first woman to win a Nobel Prize and the only person to win Nobel Prizes in two different sciences (Physics and Chemistry).
Curie discovered the elements polonium and radium. She founded the Curie Institutes in Paris and Warsaw.`,
},
{
title: "Nobel Prize in Physics",
content: `The Nobel Prize in Physics is awarded annually by the Royal Swedish Academy of Sciences.
Notable recipients include Albert Einstein (1921) for the photoelectric effect,
Marie Curie (1903) for research on radiation phenomena,
and Niels Bohr (1922) for his contributions to understanding atomic structure.`,
},
{
title: "Quantum Mechanics Foundations",
content: `Quantum mechanics emerged in the early 20th century through the work of many physicists.
Max Planck introduced the concept of energy quanta in 1900.
Niels Bohr proposed the Bohr model of the atom.
Werner Heisenberg developed the uncertainty principle.
These discoveries built upon Einstein's work on the photoelectric effect.`,
},
}
t.Log("Adding Science documents...")
for _, doc := range scienceDocs {
docID := addFixedTestDocument(t, ctx, SearchTestScienceCollection, doc.title, doc.content)
if docID != "" {
t.Logf(" ✓ Added: %s", doc.title)
}
}
// Add Tech Documents
// Entity relationships: Apple -> Steve Jobs -> iPhone -> iOS
techDocs := []struct {
title string
content string
}{
{
title: "Apple Inc History",
content: `Apple Inc. was founded by Steve Jobs, Steve Wozniak, and Ronald Wayne in 1976.
The company revolutionized personal computing with the Macintosh in 1984.
Under Steve Jobs' leadership, Apple introduced the iPhone in 2007, which transformed the smartphone industry.
Apple is headquartered in Cupertino, California.`,
},
{
title: "iPhone Development",
content: `The iPhone was introduced by Steve Jobs at Macworld 2007.
It combined a mobile phone, widescreen iPod, and internet device into one product.
The iPhone runs on iOS, Apple's mobile operating system.
The App Store, launched in 2008, created a new ecosystem for mobile applications.`,
},
{
title: "Google and AI",
content: `Google has been a pioneer in artificial intelligence and machine learning.
The company developed TensorFlow, an open-source machine learning framework.
Google's AI research includes natural language processing, computer vision, and deep learning.
Google Brain and DeepMind are the company's main AI research divisions.`,
},
{
title: "Machine Learning Applications",
content: `Machine learning is transforming various industries through AI applications.
Google uses ML for search ranking, language translation, and image recognition.
TensorFlow enables developers to build and train neural networks.
Deep learning models can now understand natural language and generate human-like text.`,
},
{
title: "Tech Industry Leaders",
content: `The technology industry has been shaped by visionary leaders.
Steve Jobs transformed Apple into the world's most valuable company.
Larry Page and Sergey Brin founded Google and pioneered internet search.
Elon Musk leads Tesla and SpaceX, pushing boundaries in electric vehicles and space exploration.`,
},
}
t.Log("Adding Tech documents...")
for _, doc := range techDocs {
docID := addFixedTestDocument(t, ctx, SearchTestTechCollection, doc.title, doc.content)
if docID != "" {
t.Logf(" ✓ Added: %s", doc.title)
}
}
// Wait for indexing
t.Log("Waiting for indexing...")
time.Sleep(2 * time.Second)
// Verify setup
t.Log("Verifying setup...")
scienceDocsResult, _ := kb.API.ListDocuments(ctx, &api.ListDocumentsFilter{
Page: 1,
PageSize: 20,
CollectionID: SearchTestScienceCollection,
})
techDocsResult, _ := kb.API.ListDocuments(ctx, &api.ListDocumentsFilter{
Page: 1,
PageSize: 20,
CollectionID: SearchTestTechCollection,
})
t.Logf("✓ Setup complete!")
t.Logf(" Science collection: %d documents", len(scienceDocsResult.Data))
t.Logf(" Tech collection: %d documents", len(techDocsResult.Data))
t.Logf("")
t.Logf("Now run search tests with:")
t.Logf(" go test -v -run 'TestSearchQuery' ./kb/api/...")
}
// ========== Cleanup Test ==========
// TestSearchCleanup removes test collections.
// Run this to clean up test data:
//
// go test -v -run "TestSearchCleanup" ./kb/api/...
func TestSearchCleanup(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
t.Log("Removing test collections...")
result1, err := kb.API.RemoveCollection(ctx, SearchTestScienceCollection)
if err != nil {
t.Logf(" Science collection removal: %v", err)
} else if result1.Removed {
t.Logf("✓ Removed: %s", SearchTestScienceCollection)
}
result2, err := kb.API.RemoveCollection(ctx, SearchTestTechCollection)
if err != nil {
t.Logf(" Tech collection removal: %v", err)
} else if result2.Removed {
t.Logf("✓ Removed: %s", SearchTestTechCollection)
}
t.Log("✓ Cleanup complete!")
}
// ========== Verify Test ==========
// TestSearchVerify checks if test collections exist and have documents.
// Run this to verify test data:
//
// go test -v -run "TestSearchVerify" ./kb/api/...
func TestSearchVerify(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
// Check Science collection
scienceExists, err := kb.API.CollectionExists(ctx, SearchTestScienceCollection)
if err != nil {
t.Fatalf("Failed to check science collection: %v", err)
}
if !scienceExists.Exists {
t.Fatalf("✗ Science collection does not exist. Run TestSearchSetup first.")
}
scienceDocs, err := kb.API.ListDocuments(ctx, &api.ListDocumentsFilter{
Page: 1,
PageSize: 20,
CollectionID: SearchTestScienceCollection,
})
assert.NoError(t, err)
t.Logf("✓ Science collection: %s (%d documents)", SearchTestScienceCollection, len(scienceDocs.Data))
for _, doc := range scienceDocs.Data {
t.Logf(" - %s", doc["name"])
}
// Check Tech collection
techExists, err := kb.API.CollectionExists(ctx, SearchTestTechCollection)
if err != nil {
t.Fatalf("Failed to check tech collection: %v", err)
}
if !techExists.Exists {
t.Fatalf("✗ Tech collection does not exist. Run TestSearchSetup first.")
}
techDocs, err := kb.API.ListDocuments(ctx, &api.ListDocumentsFilter{
Page: 1,
PageSize: 20,
CollectionID: SearchTestTechCollection,
})
assert.NoError(t, err)
t.Logf("✓ Tech collection: %s (%d documents)", SearchTestTechCollection, len(techDocs.Data))
for _, doc := range techDocs.Data {
t.Logf(" - %s", doc["name"])
}
t.Log("")
t.Log("✓ Test data verified! Ready for search tests.")
}
// ========== Helper Functions ==========
// addFixedTestDocument adds a document for search testing
func addFixedTestDocument(t *testing.T, ctx context.Context, collectionID, title, content string) string {
params := &api.AddTextParams{
CollectionID: collectionID,
Text: content,
DocID: fmt.Sprintf("%s__%s", collectionID, sanitizeTitle(title)),
Metadata: map[string]interface{}{
"title": title,
},
Chunking: &api.ProviderConfigParams{
ProviderID: "__yao.structured",
OptionID: "standard",
},
Embedding: &api.ProviderConfigParams{
ProviderID: "__yao.openai",
OptionID: "text-embedding-3-small",
},
// Enable extraction for graph-based search
Extraction: &api.ProviderConfigParams{
ProviderID: "__yao.openai",
OptionID: "gpt-4o-mini",
},
}
result, err := kb.API.AddText(ctx, params)
if err != nil {
t.Logf("Warning: Failed to add document '%s': %v", title, err)
return ""
}
return result.DocID
}
// sanitizeTitle converts title to a safe ID format
func sanitizeTitle(title string) string {
result := ""
for _, c := range title {
if (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9') {
result += string(c)
} else if c == ' ' {
result += "_"
}
}
return result
}

446
kb/api/search_test.go Normal file
View file

@ -0,0 +1,446 @@
package api_test
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
graphragtypes "github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/yao/kb"
"github.com/yaoapp/yao/kb/api"
)
// Note: TestMain is defined in collection_test.go
// Note: Test data setup is in search_setup_test.go
// ========== Search Query Tests ==========
// These tests use fixed collection IDs from search_setup_test.go
// Run TestSearchSetup first to create test data, then run these tests.
//
// Usage:
// 1. Setup (run once): go test -v -run "TestSearchSetup" ./kb/api/...
// 2. Run tests: go test -v -run "TestSearchQuery" ./kb/api/...
// 3. Cleanup (optional): go test -v -run "TestSearchCleanup" ./kb/api/...
// verifyTestDataExists checks if test collections exist, skips if not
func verifyTestDataExists(t *testing.T, ctx context.Context) {
scienceExists, _ := kb.API.CollectionExists(ctx, SearchTestScienceCollection)
techExists, _ := kb.API.CollectionExists(ctx, SearchTestTechCollection)
if scienceExists == nil || !scienceExists.Exists || techExists == nil || !techExists.Exists {
t.Skip("Test data not found. Run 'go test -v -run TestSearchSetup ./kb/api/...' first")
}
}
func TestSearchQuery(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
verifyTestDataExists(t, ctx)
t.Run("VectorSearch_SingleCollection", func(t *testing.T) {
// Test: Simple vector search in science collection
// Query about Einstein should find Einstein-related documents
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
Input: "Who is Albert Einstein and what did he discover?",
Mode: api.SearchModeVector,
PageSize: 5,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Search error (may be expected if not implemented): %v", err)
return
}
if result == nil {
t.Skip("Search not implemented yet (returned nil)")
}
assert.Greater(t, len(result.Segments), 0, "Should find segments about Einstein")
t.Logf("Vector search returned %d segments", len(result.Segments))
// Verify relevance - top results should mention Einstein
for i, seg := range result.Segments {
t.Logf(" Segment %d (score: %.4f): %s...", i, seg.Score, truncateText(seg.Text, 100))
}
})
t.Run("VectorSearch_MultipleQueries", func(t *testing.T) {
// Test: Multiple queries in same collection, results should be merged
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
Input: "relativity theory",
Mode: api.SearchModeVector,
PageSize: 3,
},
{
CollectionID: SearchTestScienceCollection,
Input: "Nobel Prize physics",
Mode: api.SearchModeVector,
PageSize: 3,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Search error: %v", err)
return
}
assert.NotNil(t, result)
t.Logf("Multi-query search returned %d merged segments", len(result.Segments))
})
t.Run("VectorSearch_CrossCollection", func(t *testing.T) {
// Test: Search across both collections
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
Input: "innovation and discovery",
Mode: api.SearchModeVector,
PageSize: 3,
},
{
CollectionID: SearchTestTechCollection,
Input: "technology innovation",
Mode: api.SearchModeVector,
PageSize: 3,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Search error: %v", err)
return
}
assert.NotNil(t, result)
t.Logf("Cross-collection search returned %d segments", len(result.Segments))
})
t.Run("ExpandSearch_EntityExpansion", func(t *testing.T) {
// Test: Expand mode should find related entities through graph
// Query: "photoelectric effect" should expand to find:
// - Einstein (discovered it)
// - Nobel Prize (awarded for it)
// - Quantum mechanics (built upon it)
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
Input: "photoelectric effect",
Mode: api.SearchModeExpand,
MaxDepth: 2,
PageSize: 5,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Expand search error: %v", err)
return
}
assert.NotNil(t, result)
t.Logf("Expand search returned %d segments", len(result.Segments))
// Check if graph data is returned
if result.Graph != nil {
t.Logf(" Graph nodes: %d, relationships: %d",
len(result.Graph.Nodes), len(result.Graph.Relationships))
}
// Verify expanded results include related entities
for i, seg := range result.Segments {
t.Logf(" Segment %d (score: %.4f): %s...", i, seg.Score, truncateText(seg.Text, 100))
}
})
t.Run("ExpandSearch_DeepAssociation", func(t *testing.T) {
// Test: Deep association through entity relationships
// Query: "Germany physics" should expand to find:
// - Einstein (born in Germany, physicist)
// - Relativity (Einstein's theory)
// - Planck (German physicist, quantum theory)
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
Input: "German physicist contributions",
Mode: api.SearchModeExpand,
MaxDepth: 3,
PageSize: 5,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Deep expand search error: %v", err)
return
}
assert.NotNil(t, result)
t.Logf("Deep expand search returned %d segments", len(result.Segments))
})
t.Run("GraphSearch_EntityTraversal", func(t *testing.T) {
// Test: Pure graph search - find segments through entity relationships
queries := []api.Query{
{
CollectionID: SearchTestTechCollection,
Input: "Steve Jobs",
Mode: api.SearchModeGraph,
MaxDepth: 2,
PageSize: 5,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Graph search error: %v", err)
return
}
assert.NotNil(t, result)
t.Logf("Graph search returned %d segments", len(result.Segments))
if result.Graph != nil {
t.Logf(" Found %d nodes, %d relationships",
len(result.Graph.Nodes), len(result.Graph.Relationships))
for _, node := range result.Graph.Nodes {
t.Logf(" Node: %s (%s)", node.ID, node.EntityType)
}
}
})
t.Run("Search_WithMessages", func(t *testing.T) {
// Test: Search using conversation history instead of direct input
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
Messages: []graphragtypes.ChatMessage{
{Role: "user", Content: "Tell me about famous physicists"},
{Role: "assistant", Content: "There are many famous physicists throughout history..."},
{Role: "user", Content: "What about Einstein specifically?"},
},
Mode: api.SearchModeVector,
PageSize: 5,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Message-based search error: %v", err)
return
}
assert.NotNil(t, result)
t.Logf("Message-based search returned %d segments", len(result.Segments))
})
t.Run("Search_WithDocumentFilter", func(t *testing.T) {
// Test: Search within a specific document
// First, get a document ID
filter := &api.ListDocumentsFilter{
Page: 1,
PageSize: 1,
CollectionID: SearchTestScienceCollection,
}
listResult, err := kb.API.ListDocuments(ctx, filter)
if err != nil || len(listResult.Data) == 0 {
t.Skip("No documents available for filter test")
}
docID, ok := listResult.Data[0]["document_id"].(string)
if !ok {
t.Skip("Could not get document ID")
}
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
DocumentID: docID,
Input: "physics discovery",
Mode: api.SearchModeVector,
PageSize: 5,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Document-filtered search error: %v", err)
return
}
assert.NotNil(t, result)
t.Logf("Document-filtered search returned %d segments", len(result.Segments))
// Verify all results are from the specified document
for _, seg := range result.Segments {
if seg.DocumentID != "" {
assert.Equal(t, docID, seg.DocumentID, "All segments should be from filtered document")
}
}
})
t.Run("Search_WithPagination", func(t *testing.T) {
// Test: Pagination
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
Input: "physics",
Mode: api.SearchModeVector,
Page: 1,
PageSize: 2,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Paginated search error: %v", err)
return
}
assert.NotNil(t, result)
assert.LessOrEqual(t, len(result.Segments), 2, "Should respect page size")
t.Logf("Page 1: %d segments, Total: %d, TotalPages: %d",
len(result.Segments), result.Total, result.TotalPages)
// Get page 2
queries[0].Page = 2
result2, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Page 2 search error: %v", err)
return
}
if result2 != nil && len(result2.Segments) > 0 {
t.Logf("Page 2: %d segments", len(result2.Segments))
}
})
t.Run("Search_WithMinScore", func(t *testing.T) {
// Test: Filter by minimum score
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
Input: "Einstein relativity",
Mode: api.SearchModeVector,
MinScore: 0.5,
PageSize: 10,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("MinScore search error: %v", err)
return
}
assert.NotNil(t, result)
t.Logf("MinScore search returned %d segments", len(result.Segments))
// Verify all results meet minimum score
for _, seg := range result.Segments {
assert.GreaterOrEqual(t, seg.Score, 0.5, "All segments should meet minimum score")
}
})
t.Run("Search_WithMetadataFilter", func(t *testing.T) {
// Test: Filter by metadata
queries := []api.Query{
{
CollectionID: SearchTestScienceCollection,
Input: "physics",
Mode: api.SearchModeVector,
Metadata: map[string]interface{}{
"title": "Albert Einstein Biography",
},
PageSize: 10,
},
}
result, err := kb.API.Search(ctx, queries)
if err != nil {
t.Logf("Metadata filter search error: %v", err)
return
}
assert.NotNil(t, result)
t.Logf("Metadata-filtered search returned %d segments", len(result.Segments))
})
}
// ========== Error Handling Tests ==========
func TestSearchErrorHandling(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
t.Run("EmptyQueries", func(t *testing.T) {
result, err := kb.API.Search(ctx, []api.Query{})
// Empty queries should return empty result or error
if err != nil {
assert.Contains(t, err.Error(), "required")
} else {
assert.NotNil(t, result)
assert.Equal(t, 0, len(result.Segments))
}
})
t.Run("MissingCollectionID", func(t *testing.T) {
queries := []api.Query{
{
Input: "test query",
Mode: api.SearchModeVector,
},
}
_, err := kb.API.Search(ctx, queries)
assert.Error(t, err)
assert.Contains(t, err.Error(), "collection")
})
t.Run("MissingInputAndMessages", func(t *testing.T) {
queries := []api.Query{
{
CollectionID: "some_collection",
Mode: api.SearchModeVector,
},
}
_, err := kb.API.Search(ctx, queries)
assert.Error(t, err)
assert.Contains(t, err.Error(), "input")
})
t.Run("NonexistentCollection", func(t *testing.T) {
queries := []api.Query{
{
CollectionID: "nonexistent_collection_xyz",
Input: "test query",
Mode: api.SearchModeVector,
},
}
_, err := kb.API.Search(ctx, queries)
assert.Error(t, err)
})
}
// ========== Helper Functions ==========
func truncateText(text string, maxLen int) string {
if len(text) <= maxLen {
return text
}
return text[:maxLen] + "..."
}

View file

@ -195,3 +195,80 @@ type AddDocumentAsyncResult struct {
JobID string `json:"job_id" yaml:"job_id"`
DocID string `json:"doc_id" yaml:"doc_id"`
}
// ========== Search Types ==========
// SearchMode defines the search strategy
type SearchMode string
const (
// SearchModeVector performs pure vector similarity search
SearchModeVector SearchMode = "vector"
// SearchModeGraph performs graph traversal to find related segments
SearchModeGraph SearchMode = "graph"
// SearchModeExpand uses graph to expand/associate entities, then enhances vector search
// This enables deeper semantic connections through entity relationships
SearchModeExpand SearchMode = "expand"
)
// Query represents a single search query
type Query struct {
// CollectionID is the collection to search in (required)
CollectionID string `json:"collection_id" yaml:"collection_id"`
// Input is the direct search query text (e.g., LLM-summarized query)
// Either Input or Messages is required; Input takes precedence if both provided
Input string `json:"input,omitempty" yaml:"input,omitempty"`
// Messages is the conversation history for context-aware search
// The last user message is used as the query if Input is empty
Messages []types.ChatMessage `json:"messages,omitempty" yaml:"messages,omitempty"`
// Mode determines the search strategy (optional, defaults to collection config or "expand")
// - vector: pure vector similarity search
// - graph: graph traversal to find related segments
// - expand: graph-based entity expansion/association + vector search
Mode SearchMode `json:"mode,omitempty" yaml:"mode,omitempty"`
// DocumentID filters results to a specific document (optional)
DocumentID string `json:"document_id,omitempty" yaml:"document_id,omitempty"`
// MinScore filters results below this similarity threshold (optional)
MinScore float64 `json:"min_score,omitempty" yaml:"min_score,omitempty"`
// Metadata filters segments by metadata fields (optional)
Metadata map[string]interface{} `json:"metadata,omitempty" yaml:"metadata,omitempty"`
// Graph search options (used when Mode is graph or hybrid)
MaxDepth int `json:"max_depth,omitempty" yaml:"max_depth,omitempty"` // Max traversal depth (default: 2)
// Pagination options
// If not specified, returns default number of results
Page int `json:"page,omitempty" yaml:"page,omitempty"` // Page number (1-based), 0 means no pagination
PageSize int `json:"pagesize,omitempty" yaml:"pagesize,omitempty"` // Number of results per page
Cursor string `json:"cursor,omitempty" yaml:"cursor,omitempty"` // Cursor for cursor-based pagination
}
// GraphData contains graph-specific search results
type GraphData struct {
Nodes []types.GraphNode `json:"nodes,omitempty" yaml:"nodes,omitempty"`
Relationships []types.GraphRelationship `json:"relationships,omitempty" yaml:"relationships,omitempty"`
}
// SearchResult represents the merged result of search operations
type SearchResult struct {
// Segments contains the merged and deduplicated text segments with scores
Segments []types.Segment `json:"segments" yaml:"segments"`
// Graph contains merged nodes and relationships (only for graph/hybrid mode)
Graph *GraphData `json:"graph,omitempty" yaml:"graph,omitempty"`
// Pagination info
Page int `json:"page,omitempty" yaml:"page,omitempty"` // Current page number
PageSize int `json:"pagesize,omitempty" yaml:"pagesize,omitempty"` // Results per page
Total int `json:"total" yaml:"total"` // Total number of results
TotalPages int `json:"pagecnt,omitempty" yaml:"pagecnt,omitempty"` // Total pages
Next int `json:"next,omitempty" yaml:"next,omitempty"` // Next page number
Prev int `json:"prev,omitempty" yaml:"prev,omitempty"` // Previous page number
Cursor string `json:"cursor,omitempty" yaml:"cursor,omitempty"` // Cursor for next page
}