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:
parent
fa471bc4c4
commit
53db8be522
6 changed files with 1521 additions and 29 deletions
|
|
@ -274,61 +274,79 @@ func (instance *KBInstance) RemoveCollection(ctx context.Context, collectionID s
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetCollection retrieves a collection by ID
|
// 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) {
|
func (instance *KBInstance) GetCollection(ctx context.Context, collectionID string) (map[string]interface{}, error) {
|
||||||
|
|
||||||
if collectionID == "" {
|
if collectionID == "" {
|
||||||
return nil, fmt.Errorf("collection ID is required")
|
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 {
|
if err != nil {
|
||||||
// Check if it's a "not found" error
|
return nil, fmt.Errorf("collection not found")
|
||||||
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)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Convert CollectionInfo to map[string]interface{}
|
// Convert database record to result map (flatten to top level)
|
||||||
// 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)
|
|
||||||
result := make(map[string]interface{})
|
result := make(map[string]interface{})
|
||||||
result["id"] = collection.ID // Primary ID field for frontend
|
for k, v := range dbRecord {
|
||||||
result["collection_id"] = collection.ID // Alias for backward compatibility
|
result[k] = v
|
||||||
|
|
||||||
// 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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if collection.Config != nil {
|
// Set standard ID fields
|
||||||
result["config"] = collection.Config
|
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
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// CollectionExists checks if a collection exists by ID
|
// 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) {
|
func (instance *KBInstance) CollectionExists(ctx context.Context, collectionID string) (*CollectionExistsResult, error) {
|
||||||
|
|
||||||
if collectionID == "" {
|
if collectionID == "" {
|
||||||
return nil, fmt.Errorf("collection ID is required")
|
return nil, fmt.Errorf("collection ID is required")
|
||||||
}
|
}
|
||||||
|
|
||||||
exists, err := instance.GraphRag.CollectionExists(ctx, collectionID)
|
// Check database (source of truth for existence)
|
||||||
if err != nil {
|
_, dbErr := instance.Config.FindCollection(collectionID, model.QueryParam{})
|
||||||
return nil, fmt.Errorf("failed to check collection existence: %w", err)
|
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{
|
return &CollectionExistsResult{
|
||||||
CollectionID: collectionID,
|
CollectionID: collectionID,
|
||||||
Exists: exists,
|
Exists: dbExists,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -32,8 +32,8 @@ type API interface {
|
||||||
AddTextAsync(ctx context.Context, params *AddTextParams) (*AddDocumentAsyncResult, error)
|
AddTextAsync(ctx context.Context, params *AddTextParams) (*AddDocumentAsyncResult, error)
|
||||||
AddURLAsync(ctx context.Context, params *AddURLParams) (*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
|
// KBInstance holds the KB instance dependencies required by the API
|
||||||
|
|
|
||||||
551
kb/api/search.go
Normal file
551
kb/api/search.go
Normal 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
400
kb/api/search_setup_test.go
Normal 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
446
kb/api/search_test.go
Normal 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] + "..."
|
||||||
|
}
|
||||||
|
|
@ -195,3 +195,80 @@ type AddDocumentAsyncResult struct {
|
||||||
JobID string `json:"job_id" yaml:"job_id"`
|
JobID string `json:"job_id" yaml:"job_id"`
|
||||||
DocID string `json:"doc_id" yaml:"doc_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
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue