Enhance Assistant stream functionality and remove history handling

- Added permission validation in the Assistant's Stream method to ensure user authorization before processing input messages.
- Introduced conversation initialization within the Stream method to prepare the context for chat interactions.
- Removed the history.go file, which previously contained a placeholder method for handling chat history, streamlining the Assistant's codebase.
- Updated the Knowledge Base API integration in collection management, ensuring all collection operations utilize the new API structure for improved consistency and error handling.
This commit is contained in:
Max 2025-12-08 10:28:26 +08:00
parent a683452d22
commit 704c0331b1
12 changed files with 1707 additions and 385 deletions

View file

@ -22,7 +22,14 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
log.Trace("[AGENT] Stream started: assistant=%s, contextID=%s", ast.ID, ctx.ID)
defer log.Trace("[AGENT] Stream ended: assistant=%s, contextID=%s", ast.ID, ctx.ID)
// Validate user permissions
var err error
err = ast.checkPermissions(ctx)
if err != nil {
return nil, err
}
// Start stream time
streamStartTime := time.Now()
// Set up interrupt handler if interrupt controller is available
@ -65,6 +72,13 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
// Now ctx.Capabilities is set, so output adapters can use it
ast.sendAgentStreamStart(ctx, streamHandler, streamStartTime)
// Initialize chat, prepare kb collection (optional) etc.
err = ast.initializeConversation(ctx, inputMessages, opts)
if err != nil {
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
return nil, err
}
// Initialize agent trace node
agentNode := ast.initAgentTraceNode(ctx, inputMessages)

79
agent/assistant/chat.go Normal file
View file

@ -0,0 +1,79 @@
package assistant
import (
"fmt"
"github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/trace/types"
)
// WithHistory merges the input messages with chat history and traces it
// This method can be overridden or extended to implement actual history loading
func (ast *Assistant) WithHistory(ctx *context.Context, input []context.Message, agentNode types.Node, options ...*context.Options) ([]context.Message, error) {
// TODO: Implement actual history loading logic here
// For now, just simulate a check and return the input messages as is
// Simulate error check (this is where actual history loading would happen)
// if some_condition {
// ast.traceAgentFail(agentNode, err)
// return nil, err
// }
fullMessages := input
// Log the chat history
ast.traceAgentHistory(ctx, agentNode, fullMessages)
return fullMessages, nil
}
// initializeConversation initialize the conversation
func (ast *Assistant) initializeConversation(ctx *context.Context, input []context.Message, options ...*context.Options) error {
var opts *context.Options
if len(options) > 0 && options[0] != nil {
opts = options[0]
} else {
opts = &context.Options{}
}
// SKIP: History (for internal calls like title/prompt etc.)
if opts.Skip != nil && opts.Skip.History {
return nil
}
chatid := ctx.ChatID
teamid := ctx.Authorized.TeamID
userid := ctx.Authorized.UserID
fmt.Printf(">>> initializeChat: chatid=%s, teamid=%s, userid=%s\n", chatid, teamid, userid)
// Prepare kb collection (optional)
err := ast.prepareKBCollection(ctx, input, opts)
if err != nil {
return err
}
// Save chat
err = ast.saveChat(ctx, input, opts)
if err != nil {
return err
}
return nil
}
// Prepare kb collection (optional)
func (ast *Assistant) prepareKBCollection(ctx *context.Context, input []context.Message, opts *context.Options) error {
_ = ctx
_ = opts
_ = input
return nil
}
func (ast *Assistant) saveChat(ctx *context.Context, input []context.Message, opts *context.Options) error {
_ = ctx
_ = input
_ = opts
return nil
}

View file

@ -1,31 +0,0 @@
package assistant
import (
"github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/trace/types"
)
// WithHistory merges the input messages with chat history and traces it
// This method can be overridden or extended to implement actual history loading
func (ast *Assistant) WithHistory(
ctx *context.Context,
inputMessages []context.Message,
agentNode types.Node,
) ([]context.Message, error) {
// TODO: Implement actual history loading logic here
// For now, just simulate a check and return the input messages as is
// Simulate error check (this is where actual history loading would happen)
// if some_condition {
// ast.traceAgentFail(agentNode, err)
// return nil, err
// }
fullMessages := inputMessages
// Log the chat history
ast.traceAgentHistory(ctx, agentNode, fullMessages)
return fullMessages, nil
}

View file

@ -0,0 +1,14 @@
package assistant
import (
"fmt"
"github.com/yaoapp/yao/agent/context"
)
func (ast *Assistant) checkPermissions(ctx *context.Context) error {
if ctx.Authorized == nil {
return fmt.Errorf("authorized information not found")
}
return nil
}

15
kb/api/api.go Normal file
View file

@ -0,0 +1,15 @@
package api
import (
"github.com/yaoapp/gou/graphrag/types"
kbtypes "github.com/yaoapp/yao/kb/types"
)
// NewAPI creates a new API instance with the provided KB dependencies
func NewAPI(graphRag types.GraphRag, config *kbtypes.Config, providers *kbtypes.ProviderConfig) API {
return &KBInstance{
GraphRag: graphRag,
Config: config,
Providers: providers,
}
}

582
kb/api/collection.go Normal file
View file

@ -0,0 +1,582 @@
package api
import (
"context"
"fmt"
graphragtypes "github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/gou/model"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/kun/maps"
)
// CreateCollection creates a new collection with the provided parameters
func (instance *KBInstance) CreateCollection(ctx context.Context, params *CreateCollectionParams) (*CreateCollectionResult, error) {
// Basic validation (before provider settings)
if params.ID == "" {
return nil, fmt.Errorf("invalid parameters: id is required")
}
if params.EmbeddingProviderID == "" {
return nil, fmt.Errorf("invalid parameters: embedding_provider_id is required")
}
if params.EmbeddingOptionID == "" {
return nil, fmt.Errorf("invalid parameters: embedding_option_id is required")
}
// Get provider settings to resolve dimension and properties
providerSettings, err := instance.getProviderSettings(params.EmbeddingProviderID, params.EmbeddingOptionID, params.Locale)
if err != nil {
return nil, fmt.Errorf("failed to resolve provider settings: %w", err)
}
// Set dimension from provider settings
if params.Config != nil {
params.Config.Dimension = providerSettings.Dimension
}
// Validate full parameters after dimension is set
if err := validateCreateParams(params); err != nil {
return nil, fmt.Errorf("invalid parameters: %w", err)
}
// Prepare metadata
metadata := params.Metadata
if metadata == nil {
metadata = make(map[string]interface{})
}
// Add embedding information to metadata
metadata["__embedding_provider"] = params.EmbeddingProviderID
metadata["__embedding_option"] = params.EmbeddingOptionID
if providerSettings.Properties != nil {
metadata["__embedding_properties"] = providerSettings.Properties
}
if params.Locale != "" {
metadata["__locale"] = params.Locale
}
// Prepare database record
dbData := map[string]interface{}{
"collection_id": params.ID,
"name": metadata["name"],
"description": metadata["description"],
"status": "creating",
"embedding_provider_id": params.EmbeddingProviderID,
"embedding_option_id": params.EmbeddingOptionID,
"embedding_properties": providerSettings.Properties,
"locale": params.Locale,
}
// Add config options to database if provided
if params.Config != nil {
if params.Config.Distance != "" {
dbData["distance"] = params.Config.Distance
}
if params.Config.IndexType != "" {
dbData["index_type"] = params.Config.IndexType
}
if params.Config.M > 0 {
dbData["m"] = params.Config.M
}
if params.Config.EfConstruction > 0 {
dbData["ef_construction"] = params.Config.EfConstruction
}
if params.Config.EfSearch > 0 {
dbData["ef_search"] = params.Config.EfSearch
}
if params.Config.NumLists > 0 {
dbData["num_lists"] = params.Config.NumLists
}
if params.Config.NumProbes > 0 {
dbData["num_probes"] = params.Config.NumProbes
}
}
// Add share field from metadata if provided
if share, ok := metadata["share"].(string); ok {
if share == "private" || share == "team" {
dbData["share"] = share
}
}
// Merge auth scope fields
if params.AuthScope != nil {
for k, v := range params.AuthScope {
dbData[k] = v
}
}
// Create database record first
_, err = instance.Config.CreateCollection(maps.MapStrAny(dbData))
if err != nil {
return nil, fmt.Errorf("failed to save collection metadata: %w", err)
}
collectionConfig := graphragtypes.CollectionConfig{
ID: params.ID,
Metadata: metadata,
Config: params.Config,
}
// Create collection in GraphRag
collectionID, err := instance.GraphRag.CreateCollection(ctx, collectionConfig)
if err != nil {
// Rollback: remove the database record
rollbackErr := instance.Config.RemoveCollection(params.ID)
if rollbackErr != nil {
log.Error("Failed to rollback collection database record: %v", rollbackErr)
}
return nil, fmt.Errorf("failed to create collection: %w", err)
}
// Update status to active after successful creation
updateErr := instance.updateCollectionWithSync(ctx, params.ID, maps.MapStrAny{"status": "active"})
if updateErr != nil {
log.Error("Failed to update collection status to active: %v", updateErr)
}
return &CreateCollectionResult{
CollectionID: collectionID,
Message: "Collection created successfully",
}, nil
}
// RemoveCollection removes an existing collection by ID
func (instance *KBInstance) RemoveCollection(ctx context.Context, collectionID string) (*RemoveCollectionResult, error) {
if collectionID == "" {
return nil, fmt.Errorf("collection ID is required")
}
removed, err := instance.GraphRag.RemoveCollection(ctx, collectionID)
if err != nil {
return nil, fmt.Errorf("failed to remove collection: %w", err)
}
if !removed {
return nil, fmt.Errorf("collection not found or could not be removed")
}
// Remove collection and documents from database after successful GraphRag removal
documentsRemoved := 0
// Count documents in this collection
if count, err := instance.Config.DocumentCount(collectionID); err == nil {
documentsRemoved = count
}
// Remove all documents belonging to this collection
if err := instance.Config.RemoveDocumentsByCollectionID(collectionID); err != nil {
log.Error("Failed to remove documents from collection %s: %v", collectionID, err)
} else {
log.Info("Removed %d documents from collection %s", documentsRemoved, collectionID)
}
// Remove the collection itself
if err := instance.Config.RemoveCollection(collectionID); err != nil {
log.Error("Failed to remove collection from database: %v", err)
} else {
log.Info("Successfully removed collection %s and %d documents", collectionID, documentsRemoved)
}
return &RemoveCollectionResult{
CollectionID: collectionID,
Removed: removed,
DocumentsRemoved: documentsRemoved,
Message: "Collection removed successfully",
}, nil
}
// GetCollection retrieves a collection by ID
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)
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)
}
// 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)
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
}
if collection.Config != nil {
result["config"] = collection.Config
}
return result, nil
}
// CollectionExists checks if a collection exists by ID
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)
}
return &CollectionExistsResult{
CollectionID: collectionID,
Exists: exists,
}, nil
}
// ListCollections lists collections with pagination and filtering
func (instance *KBInstance) ListCollections(ctx context.Context, filter *ListCollectionsFilter) (*ListCollectionsResult, error) {
page := filter.Page
if page <= 0 {
page = DefaultPage
}
pageSize := filter.PageSize
if pageSize <= 0 {
pageSize = DefaultPageSize
} else if pageSize > MaxPageSize {
pageSize = MaxPageSize
}
// Process select fields
selectFields := filter.Select
if len(selectFields) == 0 {
selectFields = DefaultCollectionFields
} else {
// Filter valid fields
validFields := []interface{}{}
for _, field := range selectFields {
if fieldStr, ok := field.(string); ok && AvailableCollectionFields[fieldStr] {
validFields = append(validFields, field)
}
}
if len(validFields) == 0 {
selectFields = DefaultCollectionFields
} else {
selectFields = validFields
}
}
// Build query parameters
param := model.QueryParam{Select: selectFields}
// Build wheres
var wheres []model.QueryWhere
// Add auth filters
if len(filter.AuthFilters) > 0 {
wheres = append(wheres, filter.AuthFilters...)
}
// Filter by keywords (search in name and description)
if filter.Keywords != "" {
wheres = append(wheres, model.QueryWhere{
Column: "name",
Value: "%" + filter.Keywords + "%",
OP: "like",
})
wheres = append(wheres, model.QueryWhere{
Column: "description",
Value: "%" + filter.Keywords + "%",
OP: "like",
Method: "orwhere",
})
}
// Filter by status
if len(filter.Status) > 0 {
statusValues := []interface{}{}
for _, status := range filter.Status {
if status != "" {
statusValues = append(statusValues, status)
}
}
if len(statusValues) > 0 {
if len(statusValues) == 1 {
wheres = append(wheres, model.QueryWhere{
Column: "status",
Value: statusValues[0],
})
} else {
wheres = append(wheres, model.QueryWhere{
Column: "status",
Value: statusValues,
OP: "in",
})
}
}
}
// Filter by system flag
if filter.System != nil {
wheres = append(wheres, model.QueryWhere{
Column: "system",
Value: *filter.System,
})
}
// Filter by embedding_provider_id
if filter.EmbeddingProviderID != "" {
wheres = append(wheres, model.QueryWhere{
Column: "embedding_provider_id",
Value: filter.EmbeddingProviderID,
})
}
param.Wheres = wheres
// Process sort orders
orders := filter.Sort
if len(orders) == 0 {
orders = DefaultSort
} else {
// Validate sort fields
validOrders := []model.QueryOrder{}
for _, order := range orders {
if ValidCollectionSortFields[order.Column] {
validOrders = append(validOrders, order)
}
}
if len(validOrders) == 0 {
orders = DefaultSort
} else {
orders = validOrders
}
}
param.Orders = orders
// Query collections
result, err := instance.Config.SearchCollections(param, page, pageSize)
if err != nil {
return nil, fmt.Errorf("failed to search collections: %w", err)
}
// Convert maps.MapStr result to ListCollectionsResult
listResult := &ListCollectionsResult{
Page: page,
PageSize: pageSize,
Data: make([]map[string]interface{}, 0), // Initialize as empty array, not nil
}
// Extract pagination data from result
if data, ok := result["data"].([]map[string]interface{}); ok {
listResult.Data = data
} else if data, ok := result["data"].([]interface{}); ok {
// Convert []interface{} to []map[string]interface{}
converted := make([]map[string]interface{}, 0, len(data))
for _, item := range data {
if mapItem, ok := item.(map[string]interface{}); ok {
converted = append(converted, mapItem)
}
}
listResult.Data = converted
} else if data, ok := result["data"].([]maps.MapStr); ok {
// Handle []maps.MapStr type (most likely from model.Paginate)
converted := make([]map[string]interface{}, 0, len(data))
for _, item := range data {
converted = append(converted, map[string]interface{}(item))
}
listResult.Data = converted
}
if next, ok := result["next"].(int); ok {
listResult.Next = next
}
if prev, ok := result["prev"].(int); ok {
listResult.Prev = prev
}
if total, ok := result["total"].(int); ok {
listResult.Total = total
}
if pagecnt, ok := result["pagecnt"].(int); ok {
listResult.PageCnt = pagecnt
}
return listResult, nil
}
// UpdateCollectionMetadata updates the metadata of an existing collection
func (instance *KBInstance) UpdateCollectionMetadata(ctx context.Context, collectionID string, params *UpdateMetadataParams) (*UpdateMetadataResult, error) {
if collectionID == "" {
return nil, fmt.Errorf("collection ID is required")
}
if len(params.Metadata) == 0 {
return nil, fmt.Errorf("metadata is required and cannot be empty")
}
err := instance.GraphRag.UpdateCollectionMetadata(ctx, collectionID, params.Metadata)
if err != nil {
return nil, fmt.Errorf("failed to update collection metadata: %w", err)
}
// Update collection metadata in database after successful GraphRag update
// Prepare update data from metadata
updateData := maps.MapStrAny{}
if name, ok := params.Metadata["name"]; ok {
updateData["name"] = name
}
if description, ok := params.Metadata["description"]; ok {
updateData["description"] = description
}
if status, ok := params.Metadata["status"]; ok {
updateData["status"] = status
}
// Merge auth scope fields
if params.AuthScope != nil {
for k, v := range params.AuthScope {
updateData[k] = v
}
}
if len(updateData) > 0 {
// Only update database, don't sync to GraphRag again
if err := instance.Config.UpdateCollection(collectionID, updateData); err != nil {
log.Error("Failed to update collection in database: %v", err)
}
}
return &UpdateMetadataResult{
CollectionID: collectionID,
Message: "Collection metadata updated successfully",
}, nil
}
// Helper methods
// validateCreateParams validates the create collection parameters
func validateCreateParams(params *CreateCollectionParams) error {
if params.ID == "" {
return fmt.Errorf("id is required")
}
if params.EmbeddingProviderID == "" {
return fmt.Errorf("embedding_provider_id is required")
}
if params.EmbeddingOptionID == "" {
return fmt.Errorf("embedding_option_id is required")
}
// Validate CreateCollectionOptions if provided
if params.Config != nil {
if err := params.Config.Validate(); err != nil && err.Error() != "collection name cannot be empty" {
return fmt.Errorf("invalid config: %w", err)
}
}
return nil
}
// ProviderSettings represents the resolved provider configuration
type ProviderSettings struct {
Dimension int `json:"dimension"`
Connector string `json:"connector"`
Properties map[string]interface{} `json:"properties"`
}
// getProviderSettings reads and resolves provider settings by provider ID and option value
func (instance *KBInstance) getProviderSettings(providerID, optionValue, locale string) (*ProviderSettings, error) {
// Default locale to "en" if empty
if locale == "" {
locale = DefaultLocale
}
// Get the specific provider from instance
provider, err := instance.Providers.GetProvider("embedding", providerID, locale)
if err != nil {
return nil, fmt.Errorf("failed to get provider %s: %v", providerID, err)
}
// Find the target option
targetOption, found := provider.GetOption(optionValue)
if !found {
return nil, fmt.Errorf("option not found: %s for provider %s", optionValue, providerID)
}
// Extract settings from option properties
settings := &ProviderSettings{
Properties: make(map[string]interface{}),
}
// Copy all properties
if targetOption.Properties != nil {
for key, value := range targetOption.Properties {
settings.Properties[key] = value
}
}
// Extract dimension
if dim, ok := targetOption.Properties["dimensions"]; ok {
if dimInt, ok := dim.(int); ok {
settings.Dimension = dimInt
} else if dimFloat, ok := dim.(float64); ok {
settings.Dimension = int(dimFloat)
}
}
// Extract connector
if connector, ok := targetOption.Properties["connector"]; ok {
if connStr, ok := connector.(string); ok {
settings.Connector = connStr
}
}
return settings, nil
}
// updateCollectionWithSync updates collection metadata in database and syncs to GraphRag
func (instance *KBInstance) updateCollectionWithSync(ctx context.Context, collectionID string, data maps.MapStrAny) error {
// Create a copy of data for GraphRag to avoid contamination from database operations
originalData := make(maps.MapStrAny)
for k, v := range data {
originalData[k] = v
}
// Update collection in database
if err := instance.Config.UpdateCollection(collectionID, data); err != nil {
return fmt.Errorf("failed to update collection in database: %w", err)
}
// Sync to GraphRag metadata
// Convert the original (unmodified) data to map[string]interface{}
metadata := make(map[string]interface{})
for k, v := range originalData {
metadata[k] = v
}
// Update GraphRag metadata
if err := instance.GraphRag.UpdateCollectionMetadata(ctx, collectionID, metadata); err != nil {
return fmt.Errorf("failed to sync collection metadata to GraphRag: %w", err)
}
return nil
}

648
kb/api/collection_test.go Normal file
View file

@ -0,0 +1,648 @@
package api_test
import (
"context"
"fmt"
"os"
"testing"
"time"
"github.com/stretchr/testify/assert"
graphragtypes "github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/gou/model"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/kb"
"github.com/yaoapp/yao/kb/api"
"github.com/yaoapp/yao/test"
)
func TestMain(m *testing.M) {
// Setup test environment
test.Prepare(&testing.T{}, config.Conf)
defer test.Clean()
// Load knowledge base
_, err := kb.Load(config.Conf)
if err != nil {
panic("Failed to load knowledge base: " + err.Error())
}
// Run tests and exit with status code
os.Exit(m.Run())
}
func TestCreateCollection(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
testCollectionID := fmt.Sprintf("test_create_%d", time.Now().UnixNano())
// Clean up after test
defer func() {
_, _ = kb.API.RemoveCollection(ctx, testCollectionID)
}()
t.Run("CreateCollectionSuccess", func(t *testing.T) {
params := &api.CreateCollectionParams{
ID: testCollectionID,
Metadata: map[string]interface{}{
"name": "Test Collection",
"description": "Test Description",
"share": "team",
},
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Locale: "en",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
M: 16,
EfConstruction: 200,
EfSearch: 64,
// Dimension will be set automatically by the API from provider settings
},
AuthScope: map[string]interface{}{
"__yao_created_by": "test_user",
"__yao_team_id": "test_team",
},
}
result, err := kb.API.CreateCollection(ctx, params)
assert.NoError(t, err)
assert.NotNil(t, result)
if result != nil {
assert.Equal(t, testCollectionID, result.CollectionID)
assert.Contains(t, result.Message, "successfully")
t.Logf("Created collection: %s", result.CollectionID)
}
})
t.Run("CreateCollectionMissingID", func(t *testing.T) {
params := &api.CreateCollectionParams{
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
result, err := kb.API.CreateCollection(ctx, params)
assert.Error(t, err)
assert.Nil(t, result)
assert.Contains(t, err.Error(), "id is required")
})
t.Run("CreateCollectionMissingProvider", func(t *testing.T) {
params := &api.CreateCollectionParams{
ID: "test_missing_provider",
EmbeddingOptionID: "text-embedding-3-small",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
result, err := kb.API.CreateCollection(ctx, params)
assert.Error(t, err)
assert.Nil(t, result)
assert.Contains(t, err.Error(), "embedding_provider_id is required")
})
t.Run("CreateCollectionInvalidProvider", func(t *testing.T) {
params := &api.CreateCollectionParams{
ID: "test_invalid_provider",
EmbeddingProviderID: "invalid_provider",
EmbeddingOptionID: "invalid_option",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
result, err := kb.API.CreateCollection(ctx, params)
assert.Error(t, err)
assert.Nil(t, result)
assert.Contains(t, err.Error(), "provider")
})
}
func TestGetCollection(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
testCollectionID := fmt.Sprintf("test_get_%d", time.Now().UnixNano())
// Create a test collection first
params := &api.CreateCollectionParams{
ID: testCollectionID,
Metadata: map[string]interface{}{
"name": "Test Get Collection",
"description": "Test Description",
},
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Locale: "en",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
_, err := kb.API.CreateCollection(ctx, params)
assert.NoError(t, err)
// Clean up after test
defer func() {
_, _ = kb.API.RemoveCollection(ctx, testCollectionID)
}()
t.Run("GetCollectionSuccess", func(t *testing.T) {
collection, err := kb.API.GetCollection(ctx, testCollectionID)
assert.NoError(t, err)
assert.NotNil(t, collection)
// Check that both id and collection_id are present
assert.Equal(t, testCollectionID, collection["id"])
assert.Equal(t, testCollectionID, collection["collection_id"])
// Check that metadata is present
assert.NotNil(t, collection["metadata"])
metadata, ok := collection["metadata"].(map[string]interface{})
assert.True(t, ok)
assert.Equal(t, "Test Get Collection", metadata["name"])
// Check that fields are also flattened at top level
assert.Equal(t, "Test Get Collection", collection["name"])
// Check that config is present
assert.NotNil(t, collection["config"])
t.Logf("Retrieved collection: %v", collection["id"])
})
t.Run("GetCollectionNotFound", func(t *testing.T) {
collection, err := kb.API.GetCollection(ctx, "nonexistent_collection")
assert.Error(t, err)
assert.Nil(t, collection)
assert.Contains(t, err.Error(), "not found")
})
t.Run("GetCollectionEmptyID", func(t *testing.T) {
collection, err := kb.API.GetCollection(ctx, "")
assert.Error(t, err)
assert.Nil(t, collection)
assert.Contains(t, err.Error(), "required")
})
}
func TestCollectionExists(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
testCollectionID := fmt.Sprintf("test_exists_%d", time.Now().UnixNano())
// Create a test collection
params := &api.CreateCollectionParams{
ID: testCollectionID,
Metadata: map[string]interface{}{
"name": "Test Exists Collection",
},
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
_, err := kb.API.CreateCollection(ctx, params)
assert.NoError(t, err)
// Clean up after test
defer func() {
_, _ = kb.API.RemoveCollection(ctx, testCollectionID)
}()
t.Run("CollectionExistsTrue", func(t *testing.T) {
result, err := kb.API.CollectionExists(ctx, testCollectionID)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.True(t, result.Exists)
assert.Equal(t, testCollectionID, result.CollectionID)
})
t.Run("CollectionExistsFalse", func(t *testing.T) {
result, err := kb.API.CollectionExists(ctx, "nonexistent_collection")
assert.NoError(t, err)
assert.NotNil(t, result)
assert.False(t, result.Exists)
})
t.Run("CollectionExistsEmptyID", func(t *testing.T) {
result, err := kb.API.CollectionExists(ctx, "")
assert.Error(t, err)
assert.Nil(t, result)
assert.Contains(t, err.Error(), "required")
})
}
func TestRemoveCollection(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
testCollectionID := fmt.Sprintf("test_remove_%d", time.Now().UnixNano())
// Create a test collection
params := &api.CreateCollectionParams{
ID: testCollectionID,
Metadata: map[string]interface{}{
"name": "Test Remove Collection",
},
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
_, err := kb.API.CreateCollection(ctx, params)
assert.NoError(t, err)
t.Run("RemoveCollectionSuccess", func(t *testing.T) {
result, err := kb.API.RemoveCollection(ctx, testCollectionID)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.True(t, result.Removed)
assert.Equal(t, testCollectionID, result.CollectionID)
assert.Contains(t, result.Message, "successfully")
// Verify collection is removed
exists, err := kb.API.CollectionExists(ctx, testCollectionID)
assert.NoError(t, err)
assert.False(t, exists.Exists)
})
t.Run("RemoveCollectionNotFound", func(t *testing.T) {
result, err := kb.API.RemoveCollection(ctx, "nonexistent_collection")
assert.Error(t, err)
assert.Nil(t, result)
})
t.Run("RemoveCollectionEmptyID", func(t *testing.T) {
result, err := kb.API.RemoveCollection(ctx, "")
assert.Error(t, err)
assert.Nil(t, result)
assert.Contains(t, err.Error(), "required")
})
}
func TestListCollections(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
// Create multiple test collections
timestamp := time.Now().UnixNano()
testCollections := []string{
fmt.Sprintf("test_list_1_%d", timestamp),
fmt.Sprintf("test_list_2_%d", timestamp),
fmt.Sprintf("test_list_3_%d", timestamp),
}
for i, collectionID := range testCollections {
params := &api.CreateCollectionParams{
ID: collectionID,
Metadata: map[string]interface{}{
"name": "Test List Collection " + string(rune('A'+i)),
"description": "Description " + string(rune('A'+i)),
},
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
_, err := kb.API.CreateCollection(ctx, params)
assert.NoError(t, err)
}
// Clean up after test
defer func() {
for _, collectionID := range testCollections {
_, _ = kb.API.RemoveCollection(ctx, collectionID)
}
}()
t.Run("ListCollectionsDefault", func(t *testing.T) {
filter := &api.ListCollectionsFilter{
Page: 1,
PageSize: 20,
}
result, err := kb.API.ListCollections(ctx, filter)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.NotNil(t, result.Data)
assert.GreaterOrEqual(t, len(result.Data), 3) // At least our 3 test collections
assert.Equal(t, 1, result.Page)
assert.Equal(t, 20, result.PageSize)
t.Logf("Found %d collections", len(result.Data))
})
t.Run("ListCollectionsWithPagination", func(t *testing.T) {
filter := &api.ListCollectionsFilter{
Page: 1,
PageSize: 2,
}
result, err := kb.API.ListCollections(ctx, filter)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.LessOrEqual(t, len(result.Data), 2)
assert.Equal(t, 1, result.Page)
assert.Equal(t, 2, result.PageSize)
})
t.Run("ListCollectionsWithKeywords", func(t *testing.T) {
filter := &api.ListCollectionsFilter{
Page: 1,
PageSize: 20,
Keywords: "Test List Collection A",
}
result, err := kb.API.ListCollections(ctx, filter)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.GreaterOrEqual(t, len(result.Data), 1)
// Check that returned collections match the keyword
for _, item := range result.Data {
name, ok := item["name"].(string)
if ok {
assert.Contains(t, name, "Test List Collection")
}
}
})
t.Run("ListCollectionsWithStatus", func(t *testing.T) {
filter := &api.ListCollectionsFilter{
Page: 1,
PageSize: 20,
Status: []string{"active"},
}
result, err := kb.API.ListCollections(ctx, filter)
assert.NoError(t, err)
assert.NotNil(t, result)
// All collections should have status "active"
for _, item := range result.Data {
status, ok := item["status"].(string)
if ok {
assert.Equal(t, "active", status)
}
}
})
t.Run("ListCollectionsWithSort", func(t *testing.T) {
filter := &api.ListCollectionsFilter{
Page: 1,
PageSize: 20,
Sort: []model.QueryOrder{
{Column: "created_at", Option: "desc"},
},
}
result, err := kb.API.ListCollections(ctx, filter)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.GreaterOrEqual(t, len(result.Data), 3)
})
t.Run("ListCollectionsWithSelect", func(t *testing.T) {
filter := &api.ListCollectionsFilter{
Page: 1,
PageSize: 20,
Select: []interface{}{"id", "collection_id", "name", "status"},
}
result, err := kb.API.ListCollections(ctx, filter)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.GreaterOrEqual(t, len(result.Data), 3)
// Check that returned fields are limited
for _, item := range result.Data {
assert.NotNil(t, item["collection_id"])
assert.NotNil(t, item["name"])
}
})
t.Run("ListCollectionsEmptyResult", func(t *testing.T) {
filter := &api.ListCollectionsFilter{
Page: 1,
PageSize: 20,
Keywords: "nonexistent_keyword_xyz123",
}
result, err := kb.API.ListCollections(ctx, filter)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.NotNil(t, result.Data)
assert.Equal(t, 0, len(result.Data))
})
}
func TestUpdateCollectionMetadata(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
testCollectionID := fmt.Sprintf("test_update_%d", time.Now().UnixNano())
// Create a test collection
params := &api.CreateCollectionParams{
ID: testCollectionID,
Metadata: map[string]interface{}{
"name": "Original Name",
"description": "Original Description",
},
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
_, err := kb.API.CreateCollection(ctx, params)
assert.NoError(t, err)
// Clean up after test
defer func() {
_, _ = kb.API.RemoveCollection(ctx, testCollectionID)
}()
t.Run("UpdateMetadataSuccess", func(t *testing.T) {
updateParams := &api.UpdateMetadataParams{
Metadata: map[string]interface{}{
"name": "Updated Name",
"description": "Updated Description",
},
AuthScope: map[string]interface{}{
"__yao_updated_by": "test_user",
},
}
result, err := kb.API.UpdateCollectionMetadata(ctx, testCollectionID, updateParams)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.Equal(t, testCollectionID, result.CollectionID)
assert.Contains(t, result.Message, "successfully")
// Verify the update
collection, err := kb.API.GetCollection(ctx, testCollectionID)
assert.NoError(t, err)
assert.Equal(t, "Updated Name", collection["name"])
assert.Equal(t, "Updated Description", collection["description"])
})
t.Run("UpdateMetadataEmptyID", func(t *testing.T) {
updateParams := &api.UpdateMetadataParams{
Metadata: map[string]interface{}{
"name": "Updated Name",
},
}
result, err := kb.API.UpdateCollectionMetadata(ctx, "", updateParams)
assert.Error(t, err)
assert.Nil(t, result)
assert.Contains(t, err.Error(), "required")
})
t.Run("UpdateMetadataEmptyMetadata", func(t *testing.T) {
updateParams := &api.UpdateMetadataParams{
Metadata: map[string]interface{}{},
}
result, err := kb.API.UpdateCollectionMetadata(ctx, testCollectionID, updateParams)
assert.Error(t, err)
assert.Nil(t, result)
assert.Contains(t, err.Error(), "empty")
})
t.Run("UpdateMetadataNotFound", func(t *testing.T) {
updateParams := &api.UpdateMetadataParams{
Metadata: map[string]interface{}{
"name": "Updated Name",
},
}
result, err := kb.API.UpdateCollectionMetadata(ctx, "nonexistent_collection", updateParams)
assert.Error(t, err)
assert.Nil(t, result)
})
}
func TestCollectionIntegration(t *testing.T) {
if kb.API == nil {
t.Skip("KB API not initialized")
}
ctx := context.Background()
testCollectionID := fmt.Sprintf("test_integration_%d", time.Now().UnixNano())
t.Run("FullCollectionLifecycle", func(t *testing.T) {
// 1. Create Collection
createParams := &api.CreateCollectionParams{
ID: testCollectionID,
Metadata: map[string]interface{}{
"name": "Integration Test Collection",
"description": "Full lifecycle test",
"share": "team",
},
EmbeddingProviderID: "__yao.openai",
EmbeddingOptionID: "text-embedding-3-small",
Locale: "en",
Config: &graphragtypes.CreateCollectionOptions{
Distance: "cosine",
IndexType: "hnsw",
},
}
createResult, err := kb.API.CreateCollection(ctx, createParams)
assert.NoError(t, err)
assert.NotNil(t, createResult)
t.Logf("Created collection: %s", createResult.CollectionID)
// 2. Check Exists
existsResult, err := kb.API.CollectionExists(ctx, testCollectionID)
assert.NoError(t, err)
assert.True(t, existsResult.Exists)
t.Logf("Collection exists: %v", existsResult.Exists)
// 3. Get Collection
collection, err := kb.API.GetCollection(ctx, testCollectionID)
assert.NoError(t, err)
assert.Equal(t, testCollectionID, collection["id"])
assert.Equal(t, testCollectionID, collection["collection_id"])
assert.Equal(t, "Integration Test Collection", collection["name"])
t.Logf("Retrieved collection: %s", collection["name"])
// 4. Update Metadata
updateParams := &api.UpdateMetadataParams{
Metadata: map[string]interface{}{
"name": "Updated Integration Test",
"description": "Updated description",
},
}
updateResult, err := kb.API.UpdateCollectionMetadata(ctx, testCollectionID, updateParams)
assert.NoError(t, err)
assert.NotNil(t, updateResult)
t.Logf("Updated collection metadata")
// 5. Verify Update
updatedCollection, err := kb.API.GetCollection(ctx, testCollectionID)
assert.NoError(t, err)
assert.Equal(t, "Updated Integration Test", updatedCollection["name"])
t.Logf("Verified update: %s", updatedCollection["name"])
// 6. List Collections (should include our test collection)
listFilter := &api.ListCollectionsFilter{
Page: 1,
PageSize: 20,
Keywords: "Updated Integration Test",
}
listResult, err := kb.API.ListCollections(ctx, listFilter)
assert.NoError(t, err)
assert.GreaterOrEqual(t, len(listResult.Data), 1)
t.Logf("Found collection in list")
// 7. Remove Collection
removeResult, err := kb.API.RemoveCollection(ctx, testCollectionID)
assert.NoError(t, err)
assert.True(t, removeResult.Removed)
t.Logf("Removed collection: %s", removeResult.CollectionID)
// 8. Verify Removal
existsAfterRemove, err := kb.API.CollectionExists(ctx, testCollectionID)
assert.NoError(t, err)
assert.False(t, existsAfterRemove.Exists)
t.Logf("Verified removal: exists=%v", existsAfterRemove.Exists)
})
}

57
kb/api/consts.go Normal file
View file

@ -0,0 +1,57 @@
package api
import "github.com/yaoapp/gou/model"
// Collection field definitions
var (
// AvailableCollectionFields defines all available fields for security filtering
AvailableCollectionFields = map[string]bool{
"id": true, "collection_id": true, "name": true, "description": true,
"status": true, "preset": true, "public": true, "share": true, "sort": true, "cover": true,
"document_count": true, "embedding_provider_id": true, "embedding_option_id": true,
"embedding_properties": true, "locale": true, "dimension": true,
"distance_metric": true, "hnsw_m": true, "ef_construction": true,
"ef_search": true, "num_lists": true, "num_probes": true,
"created_at": true, "updated_at": true,
}
// DefaultCollectionFields defines the default compact field list
DefaultCollectionFields = []interface{}{
"id", "collection_id", "name", "description", "status", "preset", "public", "share",
"sort", "cover", "document_count", "embedding_provider_id", "embedding_option_id",
"locale", "dimension", "distance_metric", "created_at", "updated_at",
}
// ValidCollectionSortFields defines valid fields for sorting
ValidCollectionSortFields = map[string]bool{
"created_at": true,
"updated_at": true,
"name": true,
"sort": true,
"document_count": true,
"status": true,
}
)
// Default pagination settings
const (
DefaultPage = 1
DefaultPageSize = 20
MaxPageSize = 100
)
// Default sort settings
const (
DefaultSortField = "created_at"
DefaultSortOrder = "desc"
)
// DefaultSort defines the default sort order for collection queries
var DefaultSort = []model.QueryOrder{
{Column: DefaultSortField, Option: DefaultSortOrder},
}
// Default locale
const (
DefaultLocale = "en"
)

34
kb/api/interfaces.go Normal file
View file

@ -0,0 +1,34 @@
package api
import (
"context"
"github.com/yaoapp/gou/graphrag/types"
kbtypes "github.com/yaoapp/yao/kb/types"
)
// API defines the unified interface for all KB operations
type API interface {
// Collection operations
CreateCollection(ctx context.Context, params *CreateCollectionParams) (*CreateCollectionResult, error)
RemoveCollection(ctx context.Context, collectionID string) (*RemoveCollectionResult, error)
GetCollection(ctx context.Context, collectionID string) (map[string]interface{}, error)
CollectionExists(ctx context.Context, collectionID string) (*CollectionExistsResult, error)
ListCollections(ctx context.Context, filter *ListCollectionsFilter) (*ListCollectionsResult, error)
UpdateCollectionMetadata(ctx context.Context, collectionID string, params *UpdateMetadataParams) (*UpdateMetadataResult, error)
// Document operations (future)
// AddDocument(ctx context.Context, params *AddDocumentParams) (*AddDocumentResult, error)
// RemoveDocument(ctx context.Context, documentID string) (*RemoveDocumentResult, error)
// ...
// Segment operations (future)
// ...
}
// KBInstance holds the KB instance dependencies required by the API
type KBInstance struct {
GraphRag types.GraphRag // GraphRag instance for vector/graph operations
Config *kbtypes.Config // KB configuration
Providers *kbtypes.ProviderConfig // Provider configurations
}

73
kb/api/types.go Normal file
View file

@ -0,0 +1,73 @@
package api
import (
"github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/gou/model"
)
// CreateCollectionParams represents the parameters for creating a collection
type CreateCollectionParams struct {
ID string `json:"id"`
Metadata map[string]interface{} `json:"metadata,omitempty"`
EmbeddingProviderID string `json:"embedding_provider_id"`
EmbeddingOptionID string `json:"embedding_option_id"`
Locale string `json:"locale,omitempty"`
Config *types.CreateCollectionOptions `json:"config,omitempty"`
AuthScope map[string]interface{} `json:"-"` // Internal: authentication scope fields
}
// CreateCollectionResult represents the result of creating a collection
type CreateCollectionResult struct {
CollectionID string `json:"collection_id"`
Message string `json:"message"`
}
// RemoveCollectionResult represents the result of removing a collection
type RemoveCollectionResult struct {
CollectionID string `json:"collection_id"`
Removed bool `json:"removed"`
DocumentsRemoved int `json:"documents_removed"`
Message string `json:"message"`
}
// CollectionExistsResult represents the result of checking if a collection exists
type CollectionExistsResult struct {
CollectionID string `json:"collection_id"`
Exists bool `json:"exists"`
}
// ListCollectionsFilter represents the filter options for listing collections
type ListCollectionsFilter struct {
Page int `json:"page"`
PageSize int `json:"pagesize"`
Keywords string `json:"keywords,omitempty"`
Status []string `json:"status,omitempty"`
System *bool `json:"system,omitempty"`
EmbeddingProviderID string `json:"embedding_provider_id,omitempty"`
Select []interface{} `json:"select,omitempty"`
Sort []model.QueryOrder `json:"sort,omitempty"`
AuthFilters []model.QueryWhere `json:"-"` // Internal: authentication filters
}
// ListCollectionsResult represents the result of listing collections
type ListCollectionsResult struct {
Data []map[string]interface{} `json:"data"`
Next int `json:"next"`
Prev int `json:"prev"`
Page int `json:"page"`
PageSize int `json:"pagesize"`
Total int `json:"total"`
PageCnt int `json:"pagecnt"`
}
// UpdateMetadataParams represents the parameters for updating collection metadata
type UpdateMetadataParams struct {
Metadata map[string]interface{} `json:"metadata"`
AuthScope map[string]interface{} `json:"-"` // Internal: authentication scope fields for update
}
// UpdateMetadataResult represents the result of updating collection metadata
type UpdateMetadataResult struct {
CollectionID string `json:"collection_id"`
Message string `json:"message"`
}

View file

@ -10,6 +10,7 @@ import (
"github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/kb/api"
// Register the built-in providers
_ "github.com/yaoapp/yao/kb/providers"
@ -21,6 +22,9 @@ import (
// Instance is the GraphRag instance
var Instance types.GraphRag = nil
// API is the Knowledge Base API instance
var API api.API = nil
// KnowledgeBase is the Knowledge Base instance
type KnowledgeBase struct {
Config *kbtypes.Config // Knowledge Base configuration
@ -86,6 +90,10 @@ func Load(appConfig config.Config) (*KnowledgeBase, error) {
// Set the instance to the global variable
Instance = instance
// Create and set the API instance
API = api.NewAPI(graphRag, &config, providers)
return instance, nil
}

View file

@ -9,9 +9,9 @@ import (
"github.com/gin-gonic/gin"
"github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/gou/model"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/kun/maps"
"github.com/yaoapp/yao/kb"
kbapi "github.com/yaoapp/yao/kb/api"
"github.com/yaoapp/yao/openapi/oauth/authorized"
oauthtypes "github.com/yaoapp/yao/openapi/oauth/types"
"github.com/yaoapp/yao/openapi/response"
@ -20,37 +20,6 @@ import (
// Collection Management Handlers
// Collection field definitions
var (
// availableCollectionFields defines all available fields for security filtering
availableCollectionFields = map[string]bool{
"id": true, "collection_id": true, "name": true, "description": true,
"status": true, "preset": true, "public": true, "share": true, "sort": true, "cover": true,
"document_count": true, "embedding_provider_id": true, "embedding_option_id": true,
"embedding_properties": true, "locale": true, "dimension": true,
"distance_metric": true, "hnsw_m": true, "ef_construction": true,
"ef_search": true, "num_lists": true, "num_probes": true,
"created_at": true, "updated_at": true,
}
// defaultCollectionFields defines the default compact field list
defaultCollectionFields = []interface{}{
"id", "collection_id", "name", "description", "status", "preset", "public", "share",
"sort", "cover", "document_count", "embedding_provider_id", "embedding_option_id",
"locale", "dimension", "distance_metric", "created_at", "updated_at",
}
// validCollectionSortFields defines valid fields for sorting
validCollectionSortFields = map[string]bool{
"created_at": true,
"updated_at": true,
"name": true,
"sort": true,
"document_count": true,
"status": true,
}
)
// ProviderSettings represents the resolved provider configuration
type ProviderSettings struct {
Dimension int `json:"dimension"`
@ -61,6 +30,16 @@ type ProviderSettings struct {
// CreateCollection creates a new collection
func CreateCollection(c *gin.Context) {
// Check if kb.API is available
if kb.API == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Knowledge base not initialized",
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Prepare request and database data
req, collectionData, err := PrepareCreateCollection(c)
if err != nil {
@ -74,12 +53,56 @@ func CreateCollection(c *gin.Context) {
// Attach create scope to the collection data
authInfo := authorized.GetInfo(c)
var authScope map[string]interface{}
if authInfo != nil {
collectionData = authInfo.WithCreateScope(collectionData)
// Extract auth scope fields
authScope = make(map[string]interface{})
if createdBy, ok := collectionData["__yao_created_by"]; ok {
authScope["__yao_created_by"] = createdBy
}
if updatedBy, ok := collectionData["__yao_updated_by"]; ok {
authScope["__yao_updated_by"] = updatedBy
}
if teamID, ok := collectionData["__yao_team_id"]; ok {
authScope["__yao_team_id"] = teamID
}
}
// Check if kb.Instance is available
if kb.Instance == nil {
// Build API params
params := &kbapi.CreateCollectionParams{
ID: req.ID,
Metadata: req.Metadata,
EmbeddingProviderID: req.Config.EmbeddingProviderID,
EmbeddingOptionID: req.Config.EmbeddingOptionID,
Locale: req.Config.Locale,
Config: req.Config.CreateCollectionOptions,
AuthScope: authScope,
}
// Call API to create collection
result, err := kb.API.CreateCollection(c.Request.Context(), params)
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
successData := gin.H{
"message": result.Message,
"collection_id": result.CollectionID,
}
response.RespondWithSuccess(c, response.StatusCreated, successData)
}
// RemoveCollection removes an existing collection
func RemoveCollection(c *gin.Context) {
// Check if kb.API is available
if kb.API == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Knowledge base not initialized",
@ -88,68 +111,6 @@ func CreateCollection(c *gin.Context) {
return
}
// Get KB config
config, err := kb.GetConfig()
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to get KB config: " + err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// First create database record
_, err = config.CreateCollection(maps.MapStrAny(collectionData))
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to save collection metadata: " + err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Create CollectionConfig for GraphRag
collectionConfig := types.CollectionConfig{
ID: req.ID,
Metadata: req.Metadata,
Config: req.Config.CreateCollectionOptions,
}
// Call the actual CreateCollection method
collectionID, err := kb.Instance.CreateCollection(c.Request.Context(), collectionConfig)
if err != nil {
// Rollback: remove the database record
rollbackErr := config.RemoveCollection(req.ID)
if rollbackErr != nil {
log.Error("Failed to rollback collection database record: %v", rollbackErr)
}
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to create collection: " + err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Update status to active after successful creation and sync to GraphRag
updateErr := UpdateCollectionWithSync(req.ID, maps.MapStrAny{"status": "active"}, config)
if updateErr != nil {
log.Error("Failed to update collection status to active: %v", updateErr)
}
successData := gin.H{
"message": "Collection created successfully",
"collection_id": collectionID,
}
response.RespondWithSuccess(c, response.StatusCreated, successData)
}
// RemoveCollection removes an existing collection
func RemoveCollection(c *gin.Context) {
authInfo := authorized.GetInfo(c)
// Get collection ID from URL parameter
@ -163,16 +124,6 @@ func RemoveCollection(c *gin.Context) {
return
}
// Check if kb.Instance is available
if kb.Instance == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Knowledge base not initialized",
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Check remove permission
hasPermission, err := checkCollectionPermission(authInfo, collectionID)
if err != nil {
@ -194,60 +145,38 @@ func RemoveCollection(c *gin.Context) {
return
}
// Call the actual RemoveCollection method
removed, err := kb.Instance.RemoveCollection(c.Request.Context(), collectionID)
// Call API to remove collection
result, err := kb.API.RemoveCollection(c.Request.Context(), collectionID)
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to remove collection: " + err.Error(),
ErrorDescription: err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
if !removed {
errorResp := &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code,
ErrorDescription: "Collection not found or could not be removed",
}
response.RespondWithError(c, response.StatusNotFound, errorResp)
return
}
// Remove collection and all its documents from database after successful GraphRag removal
documentsRemoved := 0
if config, err := kb.GetConfig(); err == nil {
// First, count documents in this collection (for reporting)
if count, err := config.DocumentCount(collectionID); err == nil {
documentsRemoved = count
}
// Remove all documents belonging to this collection
if err := config.RemoveDocumentsByCollectionID(collectionID); err != nil {
log.Error("Failed to remove documents from collection %s: %v", collectionID, err)
} else {
log.Info("Removed %d documents from collection %s", documentsRemoved, collectionID)
}
// Then remove the collection itself
if err := config.RemoveCollection(collectionID); err != nil {
log.Error("Failed to remove collection from database: %v", err)
} else {
log.Info("Successfully removed collection %s and %d documents", collectionID, documentsRemoved)
}
}
successData := gin.H{
"message": "Collection removed successfully",
"collection_id": collectionID,
"removed": removed,
"documents_removed": documentsRemoved,
"message": result.Message,
"collection_id": result.CollectionID,
"removed": result.Removed,
"documents_removed": result.DocumentsRemoved,
}
response.RespondWithSuccess(c, response.StatusOK, successData)
}
// CollectionExists checks if a collection exists
func CollectionExists(c *gin.Context) {
// Check if kb.API is available
if kb.API == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Knowledge base not initialized",
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Get collection ID from URL parameter
collectionID := c.Param("collectionID")
if collectionID == "" {
@ -259,8 +188,28 @@ func CollectionExists(c *gin.Context) {
return
}
// Check if kb.Instance is available
if kb.Instance == nil {
// Call API to check collection existence
result, err := kb.API.CollectionExists(c.Request.Context(), collectionID)
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
successData := gin.H{
"collection_id": result.CollectionID,
"exists": result.Exists,
}
response.RespondWithSuccess(c, response.StatusOK, successData)
}
// GetCollection retrieves a collection by ID
func GetCollection(c *gin.Context) {
// Check if kb.API is available
if kb.API == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Knowledge base not initialized",
@ -269,26 +218,6 @@ func CollectionExists(c *gin.Context) {
return
}
// Call the actual CollectionExists method
exists, err := kb.Instance.CollectionExists(c.Request.Context(), collectionID)
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to check collection existence: " + err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
successData := gin.H{
"collection_id": collectionID,
"exists": exists,
}
response.RespondWithSuccess(c, response.StatusOK, successData)
}
// GetCollection retrieves a collection by ID
func GetCollection(c *gin.Context) {
collectionID := c.Param("collectionID")
if collectionID == "" {
errorResp := &response.ErrorResponse{
@ -299,21 +228,11 @@ func GetCollection(c *gin.Context) {
return
}
// Check if kb.Instance is available
if kb.Instance == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Knowledge base not initialized",
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Use the dedicated GetCollection method
collection, err := kb.Instance.GetCollection(c.Request.Context(), collectionID)
// Call API to get collection
collection, err := kb.API.GetCollection(c.Request.Context(), collectionID)
if err != nil {
// Check if it's a "not found" error
if err.Error() == fmt.Sprintf("collection with ID '%s' not found", collectionID) {
if err.Error() == "collection not found" || err.Error() == fmt.Sprintf("collection with ID '%s' not found", collectionID) {
errorResp := &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code,
ErrorDescription: "Collection not found",
@ -324,7 +243,7 @@ func GetCollection(c *gin.Context) {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to get collection: " + err.Error(),
ErrorDescription: err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
@ -336,11 +255,8 @@ func GetCollection(c *gin.Context) {
// ListCollections lists collections with pagination
func ListCollections(c *gin.Context) {
// Get authorized information
authInfo := authorized.GetInfo(c)
// Check if kb.Instance is available
if kb.Instance == nil {
// Check if kb.API is available
if kb.API == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Knowledge base not initialized",
@ -349,6 +265,9 @@ func ListCollections(c *gin.Context) {
return
}
// Get authorized information
authInfo := authorized.GetInfo(c)
// Parse pagination parameters
page := 1
if pageStr := c.Query("page"); pageStr != "" {
@ -364,180 +283,104 @@ func ListCollections(c *gin.Context) {
}
}
// Get KB config
config, err := kb.GetConfig()
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to get KB config: " + err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Parse select parameter
var selectFields []interface{}
if selectParam := strings.TrimSpace(c.Query("select")); selectParam != "" {
requestedFields := strings.Split(selectParam, ",")
for _, field := range requestedFields {
field = strings.TrimSpace(field)
if field != "" && availableCollectionFields[field] {
if field != "" && kbapi.AvailableCollectionFields[field] {
selectFields = append(selectFields, field)
}
}
// If no valid fields found, use default
if len(selectFields) == 0 {
selectFields = defaultCollectionFields
}
// Parse sort parameter
var orders []model.QueryOrder
if sortParam := strings.TrimSpace(c.Query("sort")); sortParam != "" {
sortItems := strings.Split(sortParam, ",")
for _, sortItem := range sortItems {
sortItem = strings.TrimSpace(sortItem)
if sortItem == "" {
continue
}
sortParts := strings.Fields(sortItem)
if len(sortParts) == 0 {
continue
}
sortField := sortParts[0]
sortOrder := "desc"
if len(sortParts) >= 2 {
sortOrder = strings.ToLower(sortParts[1])
}
// Validate sort field and order
if kbapi.ValidCollectionSortFields[sortField] && (sortOrder == "asc" || sortOrder == "desc") {
orders = append(orders, model.QueryOrder{
Column: sortField,
Option: sortOrder,
})
}
}
} else {
selectFields = defaultCollectionFields
}
// Build query parameters
param := model.QueryParam{Select: selectFields}
// Add filters
var wheres []model.QueryWhere
// Apply permission-based filtering
wheres = append(wheres, AuthFilter(c, authInfo)...)
// Filter by keywords (search in name and description)
if keywords := strings.TrimSpace(c.Query("keywords")); keywords != "" {
wheres = append(wheres, model.QueryWhere{
Column: "name",
Value: "%" + keywords + "%",
OP: "like",
})
wheres = append(wheres, model.QueryWhere{
Column: "description",
Value: "%" + keywords + "%",
OP: "like",
Wheres: []model.QueryWhere{},
Method: "orwhere",
})
// Build filter for API
filter := &kbapi.ListCollectionsFilter{
Page: page,
PageSize: pagesize,
Keywords: strings.TrimSpace(c.Query("keywords")),
EmbeddingProviderID: strings.TrimSpace(c.Query("embedding_provider_id")),
Select: selectFields,
Sort: orders,
AuthFilters: AuthFilter(c, authInfo),
}
// Filter by status (support multiple values separated by comma)
// Parse status parameter
if statusParam := strings.TrimSpace(c.Query("status")); statusParam != "" {
statusList := strings.Split(statusParam, ",")
var statusValues []interface{}
for _, status := range statusList {
status = strings.TrimSpace(status)
if status != "" {
statusValues = append(statusValues, status)
}
}
if len(statusValues) > 0 {
if len(statusValues) == 1 {
// Single status
wheres = append(wheres, model.QueryWhere{
Column: "status",
Value: statusValues[0],
})
} else {
// Multiple status - use IN clause
wheres = append(wheres, model.QueryWhere{
Column: "status",
Value: statusValues,
OP: "in",
})
filter.Status = append(filter.Status, status)
}
}
}
// Filter by system flag
// Parse system parameter
if systemParam := strings.TrimSpace(c.Query("system")); systemParam != "" {
switch systemParam {
case "true", "1":
wheres = append(wheres, model.QueryWhere{
Column: "system",
Value: true,
})
systemVal := true
filter.System = &systemVal
case "false", "0":
wheres = append(wheres, model.QueryWhere{
Column: "system",
Value: false,
})
systemVal := false
filter.System = &systemVal
}
}
// Filter by embedding_provider_id
if providerID := strings.TrimSpace(c.Query("embedding_provider_id")); providerID != "" {
wheres = append(wheres, model.QueryWhere{
Column: "embedding_provider_id",
Value: providerID,
})
}
param.Wheres = wheres
// Add ordering
sortParam := strings.TrimSpace(c.Query("sort"))
if sortParam == "" {
sortParam = "created_at desc" // Default sort
}
// Parse sort parameter (format: "field1 direction1,field2 direction2")
var orders []model.QueryOrder
sortItems := strings.Split(sortParam, ",")
for _, sortItem := range sortItems {
sortItem = strings.TrimSpace(sortItem)
if sortItem == "" {
continue
}
// Parse each sort item (format: "field direction")
sortParts := strings.Fields(sortItem)
sortField := "created_at" // Default field
sortOrder := "desc" // Default order
if len(sortParts) >= 1 {
sortField = sortParts[0]
}
if len(sortParts) >= 2 {
sortOrder = strings.ToLower(sortParts[1])
}
// Validate sort field
if !validCollectionSortFields[sortField] {
continue // Skip invalid fields
}
// Validate sort order
if sortOrder != "asc" && sortOrder != "desc" {
sortOrder = "desc" // Default order
}
orders = append(orders, model.QueryOrder{
Column: sortField,
Option: sortOrder,
})
}
// If no valid orders found, use default
if len(orders) == 0 {
orders = []model.QueryOrder{
{Column: "created_at", Option: "desc"},
}
}
param.Orders = orders
// Query collections using KB config
result, err := config.SearchCollections(param, page, pagesize)
// Call API to list collections
result, err := kb.API.ListCollections(c.Request.Context(), filter)
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to search collections: " + err.Error(),
ErrorDescription: err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
c.JSON(http.StatusOK, result)
// Return the result directly to maintain backward compatibility
c.JSON(http.StatusOK, gin.H{
"data": result.Data,
"next": result.Next,
"prev": result.Prev,
"page": result.Page,
"pagesize": result.PageSize,
"total": result.Total,
"pagecnt": result.PageCnt,
})
}
// UpdateCollectionMetadata updates the metadata of an existing collection
@ -576,8 +419,8 @@ func UpdateCollectionMetadata(c *gin.Context) {
return
}
// Check if kb.Instance is available
if kb.Instance == nil {
// Check if kb.API is available
if kb.API == nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Knowledge base not initialized",
@ -608,45 +451,31 @@ func UpdateCollectionMetadata(c *gin.Context) {
return
}
// Call the actual UpdateCollectionMetadata method
err = kb.Instance.UpdateCollectionMetadata(c.Request.Context(), collectionID, req.Metadata)
// Build API params
var authScope map[string]interface{}
if authInfo != nil {
authScope = authInfo.WithUpdateScope(maps.MapStrAny{})
}
params := &kbapi.UpdateMetadataParams{
Metadata: req.Metadata,
AuthScope: authScope,
}
// Call API to update collection metadata
result, err := kb.API.UpdateCollectionMetadata(c.Request.Context(), collectionID, params)
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code,
ErrorDescription: "Failed to update collection metadata: " + err.Error(),
ErrorDescription: err.Error(),
}
response.RespondWithError(c, response.StatusInternalServerError, errorResp)
return
}
// Update collection metadata in database after successful GraphRag update
// Note: Only update database here, don't sync to GraphRag again (already done above)
if config, err := kb.GetConfig(); err == nil {
// Prepare update data from metadata
updateData := maps.MapStrAny{}
if name, ok := req.Metadata["name"]; ok {
updateData["name"] = name
}
if description, ok := req.Metadata["description"]; ok {
updateData["description"] = description
}
if status, ok := req.Metadata["status"]; ok {
updateData["status"] = status
}
// Update __yao_updated_by
updateData = authInfo.WithUpdateScope(updateData)
if len(updateData) > 0 {
// Only update database, don't sync to GraphRag again to avoid duplicate updates
if err := config.UpdateCollection(collectionID, updateData); err != nil {
log.Error("Failed to update collection in database: %v", err)
}
}
}
successData := gin.H{
"message": "Collection metadata updated successfully",
"collection_id": collectionID,
"message": result.Message,
"collection_id": result.CollectionID,
}
response.RespondWithSuccess(c, response.StatusOK, successData)
}