Refactor Authorization Functions and Update Tests

- Renamed functions for consistency and clarity, changing `buildDBAuthWheres` to `BuildDBAuthWheres` and `filterKBCollectionsByAuth` to `FilterKBCollectionsByAuth`.
- Enhanced test cases to utilize the updated function names, ensuring proper authorization checks in the search functionality.
- Improved test environment initialization to streamline setup processes and ensure robust testing of authorization logic.
- Verified that search results adhere to authorization constraints, ensuring only accessible collections are queried based on user permissions.
This commit is contained in:
Max 2025-12-20 15:02:01 +08:00
parent b63f3fe246
commit e46645a986
4 changed files with 242 additions and 100 deletions

View file

@ -803,7 +803,7 @@ func (ast *Assistant) buildSearchRequests(ctx *context.Context, query string, co
}
// Filter collections by authorization (Collection-level permission check)
allowedCollections := filterKBCollectionsByAuth(ctx, ast.KB.Collections)
allowedCollections := FilterKBCollectionsByAuth(ctx, ast.KB.Collections)
if len(allowedCollections) == 0 {
ctx.Logger.Info("No accessible KB collections after auth filter")
} else {
@ -839,7 +839,7 @@ func (ast *Assistant) buildSearchRequests(ctx *context.Context, query string, co
}
// Apply authorization where clauses
if authWheres := buildDBAuthWheres(ctx); authWheres != nil {
if authWheres := BuildDBAuthWheres(ctx); authWheres != nil {
dbReq.Wheres = authWheres
}

View file

@ -5,10 +5,10 @@ import (
"github.com/yaoapp/yao/agent/context"
)
// buildDBAuthWheres builds where clauses for DB search based on authorization
// BuildDBAuthWheres builds where clauses for DB search based on authorization
// This applies permission-based filtering to database queries
// Returns gou.Where clauses to filter records by authorization scope
func buildDBAuthWheres(ctx *context.Context) []gou.Where {
func BuildDBAuthWheres(ctx *context.Context) []gou.Where {
if ctx == nil || ctx.Authorized == nil {
return nil
}

View file

@ -1,14 +1,16 @@
package assistant
package assistant_test
import (
"context"
"fmt"
"os"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
graphragtypes "github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/yao/agent/assistant"
agentContext "github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/agent/search"
searchTypes "github.com/yaoapp/yao/agent/search/types"
@ -35,23 +37,39 @@ const (
TestTeam2 = "team_2"
)
// ========== Test Environment ==========
var (
testEnvOnce sync.Once
testEnvErr error
)
// initTestEnv initializes the test environment (only once)
func initTestEnv(t *testing.T) {
testEnvOnce.Do(func() {
// Setup test environment
test.Prepare(t, config.Conf)
// Load attachment managers
if err := attachment.Load(config.Conf); err != nil {
t.Logf("Warning: Failed to load attachment managers: %v", err)
}
// Load knowledge base
if _, err := kb.Load(config.Conf); err != nil {
testEnvErr = fmt.Errorf("failed to load knowledge base: %w", err)
return
}
})
if testEnvErr != nil {
t.Fatalf("Test environment initialization failed: %v", testEnvErr)
}
}
// ========== TestMain ==========
func TestMain(m *testing.M) {
// Setup test environment
test.Prepare(&testing.T{}, config.Conf)
defer test.Clean()
// Load attachment managers
if err := attachment.Load(config.Conf); err != nil {
fmt.Printf("Warning: Failed to load attachment managers: %v\n", err)
}
// Load knowledge base
if _, err := kb.Load(config.Conf); err != nil {
fmt.Printf("Warning: Failed to load knowledge base: %v\n", err)
}
os.Exit(m.Run())
}
@ -62,6 +80,8 @@ func TestMain(m *testing.M) {
//
// go test -v -run "TestAuthSearchSetup" ./agent/assistant/...
func TestAuthSearchSetup(t *testing.T) {
initTestEnv(t)
if kb.API == nil {
t.Fatal("KB API not initialized")
}
@ -111,6 +131,8 @@ func TestAuthSearchSetup(t *testing.T) {
// TestAuthSearchCleanup removes auth test collections.
func TestAuthSearchCleanup(t *testing.T) {
initTestEnv(t)
if kb.API == nil {
t.Fatal("KB API not initialized")
}
@ -124,9 +146,11 @@ func TestAuthSearchCleanup(t *testing.T) {
// Note: KB permission filtering works at the Collection level.
// The Collection metadata contains __yao_team_id, __yao_created_by, public, share fields.
// filterKBCollectionsByAuth filters collections based on user authorization.
// FilterKBCollectionsByAuth filters collections based on user authorization.
func TestKBCollectionAuthFilter(t *testing.T) {
initTestEnv(t)
if kb.API == nil {
t.Fatal("KB API not initialized")
}
@ -139,7 +163,7 @@ func TestKBCollectionAuthFilter(t *testing.T) {
ctx := createAuthContext(TestUserA, TestTeam1, true, false)
collections := []string{AuthTestCollectionTeam1, AuthTestCollectionTeam2}
allowed := filterKBCollectionsByAuth(ctx, collections)
allowed := assistant.FilterKBCollectionsByAuth(ctx, collections)
assert.Contains(t, allowed, AuthTestCollectionTeam1, "Team1 member should access Team1 collection")
t.Logf(" Allowed collections: %v", allowed)
})
@ -149,7 +173,7 @@ func TestKBCollectionAuthFilter(t *testing.T) {
ctx := createAuthContext(TestUserA, TestTeam1, true, false)
collections := []string{AuthTestCollectionTeam2}
allowed := filterKBCollectionsByAuth(ctx, collections)
allowed := assistant.FilterKBCollectionsByAuth(ctx, collections)
assert.NotContains(t, allowed, AuthTestCollectionTeam2, "Team1 member should NOT access Team2 collection")
t.Logf(" Allowed collections: %v (expected empty)", allowed)
})
@ -159,7 +183,7 @@ func TestKBCollectionAuthFilter(t *testing.T) {
ctx := createAuthContext(TestUserA, "", false, true)
collections := []string{AuthTestCollectionTeam1, AuthTestCollectionTeam2}
allowed := filterKBCollectionsByAuth(ctx, collections)
allowed := assistant.FilterKBCollectionsByAuth(ctx, collections)
assert.Contains(t, allowed, AuthTestCollectionTeam1, "Owner should access own collection")
assert.NotContains(t, allowed, AuthTestCollectionTeam2, "Owner should NOT access other's collection")
t.Logf(" Allowed collections: %v", allowed)
@ -184,7 +208,7 @@ func TestKBCollectionAuthFilter(t *testing.T) {
ctx := createAuthContext(TestUserA, TestTeam1, false, true) // Owner check
collections := []string{AuthTestCollectionPublic}
allowed := filterKBCollectionsByAuth(ctx, collections)
allowed := assistant.FilterKBCollectionsByAuth(ctx, collections)
assert.Contains(t, allowed, AuthTestCollectionPublic, "Owner should access their collection")
t.Logf(" Allowed collections (owner check): %v", allowed)
})
@ -194,7 +218,7 @@ func TestKBCollectionAuthFilter(t *testing.T) {
ctx := createAuthContext(TestUserA, TestTeam1, false, false)
collections := []string{AuthTestCollectionTeam1, AuthTestCollectionTeam2, AuthTestCollectionPublic}
allowed := filterKBCollectionsByAuth(ctx, collections)
allowed := assistant.FilterKBCollectionsByAuth(ctx, collections)
assert.Len(t, allowed, 3, "No constraints should allow all collections")
t.Logf(" Allowed collections: %v", allowed)
})
@ -202,7 +226,7 @@ func TestKBCollectionAuthFilter(t *testing.T) {
t.Run("NilContextMeansFullAccess", func(t *testing.T) {
collections := []string{AuthTestCollectionTeam1, AuthTestCollectionTeam2}
allowed := filterKBCollectionsByAuth(nil, collections)
allowed := assistant.FilterKBCollectionsByAuth(nil, collections)
assert.Len(t, allowed, 2, "Nil context should allow all collections")
t.Logf(" Allowed collections: %v", allowed)
})
@ -213,58 +237,123 @@ func TestKBCollectionAuthFilter(t *testing.T) {
func TestDBAuthWheresFilter(t *testing.T) {
t.Run("TeamOnlyGeneratesCorrectWheres", func(t *testing.T) {
ctx := createAuthContext(TestUserA, TestTeam1, true, false)
wheres := buildDBAuthWheres(ctx)
wheres := assistant.BuildDBAuthWheres(ctx)
assert.NotNil(t, wheres)
assert.Len(t, wheres, 1)
// Verify structure contains team filter
// Verify structure: should have 2 top-level conditions (public OR team filter)
where := wheres[0]
assert.NotEmpty(t, where.Wheres)
t.Logf(" TeamOnly: Generated %d nested where clauses", len(where.Wheres))
assert.Len(t, where.Wheres, 2, "Should have 2 conditions: public OR team")
// First condition: public = true (OR)
publicCond := where.Wheres[0]
assert.NotNil(t, publicCond.Condition.Field)
assert.Equal(t, "public", publicCond.Condition.Field.Field)
assert.Equal(t, true, publicCond.Condition.Value)
assert.True(t, publicCond.Condition.OR)
// Second condition: team filter with nested conditions
teamCond := where.Wheres[1]
assert.Len(t, teamCond.Wheres, 2, "Team filter should have team_id and (created_by OR share)")
// Team ID check
teamIDCond := teamCond.Wheres[0]
assert.Equal(t, "__yao_team_id", teamIDCond.Condition.Field.Field)
assert.Equal(t, TestTeam1, teamIDCond.Condition.Value)
// Created by OR share = team
ownerOrShareCond := teamCond.Wheres[1]
assert.Len(t, ownerOrShareCond.Wheres, 2)
assert.Equal(t, "__yao_created_by", ownerOrShareCond.Wheres[0].Condition.Field.Field)
assert.Equal(t, TestUserA, ownerOrShareCond.Wheres[0].Condition.Value)
assert.Equal(t, "share", ownerOrShareCond.Wheres[1].Condition.Field.Field)
assert.Equal(t, "team", ownerOrShareCond.Wheres[1].Condition.Value)
assert.True(t, ownerOrShareCond.Wheres[1].Condition.OR)
t.Logf(" TeamOnly: Verified team_id=%s, created_by=%s", TestTeam1, TestUserA)
})
t.Run("OwnerOnlyGeneratesCorrectWheres", func(t *testing.T) {
ctx := createAuthContext(TestUserA, "", false, true)
wheres := buildDBAuthWheres(ctx)
wheres := assistant.BuildDBAuthWheres(ctx)
assert.NotNil(t, wheres)
assert.Len(t, wheres, 1)
// Verify structure contains owner filter
// Verify structure: should have 2 top-level conditions (public OR owner filter)
where := wheres[0]
assert.NotEmpty(t, where.Wheres)
t.Logf(" OwnerOnly: Generated %d nested where clauses", len(where.Wheres))
assert.Len(t, where.Wheres, 2, "Should have 2 conditions: public OR owner")
// First condition: public = true (OR)
publicCond := where.Wheres[0]
assert.NotNil(t, publicCond.Condition.Field)
assert.Equal(t, "public", publicCond.Condition.Field.Field)
assert.Equal(t, true, publicCond.Condition.Value)
assert.True(t, publicCond.Condition.OR)
// Second condition: owner filter with nested conditions
ownerCond := where.Wheres[1]
assert.Len(t, ownerCond.Wheres, 2, "Owner filter should have team_id IS NULL and created_by")
// Team ID is null check
teamNullCond := ownerCond.Wheres[0]
assert.Equal(t, "__yao_team_id", teamNullCond.Condition.Field.Field)
assert.Equal(t, "null", teamNullCond.Condition.OP)
// Created by check
createdByCond := ownerCond.Wheres[1]
assert.Equal(t, "__yao_created_by", createdByCond.Condition.Field.Field)
assert.Equal(t, TestUserA, createdByCond.Condition.Value)
t.Logf(" OwnerOnly: Verified created_by=%s, team_id IS NULL", TestUserA)
})
t.Run("NoConstraintsReturnsNil", func(t *testing.T) {
ctx := createAuthContext(TestUserA, TestTeam1, false, false)
wheres := buildDBAuthWheres(ctx)
wheres := assistant.BuildDBAuthWheres(ctx)
assert.Nil(t, wheres)
assert.Nil(t, wheres, "No constraints should return nil")
t.Log(" No constraints: nil wheres (no filter)")
})
t.Run("EmptyTeamIDReturnsNil", func(t *testing.T) {
ctx := createAuthContext(TestUserA, "", true, false)
wheres := buildDBAuthWheres(ctx)
wheres := assistant.BuildDBAuthWheres(ctx)
assert.Nil(t, wheres)
assert.Nil(t, wheres, "Empty TeamID with TeamOnly should return nil")
t.Log(" Empty TeamID with TeamOnly: nil wheres")
})
t.Run("EmptyUserIDReturnsNil", func(t *testing.T) {
ctx := createAuthContext("", TestTeam1, false, true)
wheres := buildDBAuthWheres(ctx)
wheres := assistant.BuildDBAuthWheres(ctx)
assert.Nil(t, wheres)
assert.Nil(t, wheres, "Empty UserID with OwnerOnly should return nil")
t.Log(" Empty UserID with OwnerOnly: nil wheres")
})
t.Run("NilContextReturnsNil", func(t *testing.T) {
wheres := assistant.BuildDBAuthWheres(nil)
assert.Nil(t, wheres, "Nil context should return nil")
t.Log(" Nil context: nil wheres")
})
t.Run("NilAuthorizedReturnsNil", func(t *testing.T) {
ctx := &agentContext.Context{Authorized: nil}
wheres := assistant.BuildDBAuthWheres(ctx)
assert.Nil(t, wheres, "Nil Authorized should return nil")
t.Log(" Nil Authorized: nil wheres")
})
}
// ========== KB Search Integration Tests ==========
func TestKBSearchIntegration(t *testing.T) {
initTestEnv(t)
if kb.API == nil {
t.Fatal("KB API not initialized")
}
@ -272,71 +361,117 @@ func TestKBSearchIntegration(t *testing.T) {
// Ensure test data exists
TestAuthSearchSetup(t)
t.Run("SearchWithoutFilterFindsDocuments", func(t *testing.T) {
// Search without any auth filter
result := executeKBSearch(t, AuthTestCollectionTeam1, "quantum physics machine learning", nil)
assert.Greater(t, len(result.Items), 0, "Should find documents without filter")
t.Logf(" Found %d items without filter", len(result.Items))
})
t.Run("SearchPublicCollectionWorks", func(t *testing.T) {
// Public collection should be accessible
result := executeKBSearch(t, AuthTestCollectionPublic, "artificial intelligence robotics", nil)
assert.Greater(t, len(result.Items), 0, "Public collection should be searchable")
t.Logf(" Found %d items in public collection", len(result.Items))
})
t.Run("SearchWithMetadataFilterWorks", func(t *testing.T) {
// Search with collection_id filter (this exists in segment metadata)
metadata := map[string]interface{}{
"collection_id": AuthTestCollectionTeam1,
}
result := executeKBSearch(t, AuthTestCollectionTeam1, "quantum", metadata)
t.Logf(" Found %d items with collection_id filter", len(result.Items))
// Verify all results have correct collection_id
for _, item := range result.Items {
if item.Metadata != nil {
collID, _ := item.Metadata["collection_id"].(string)
assert.Equal(t, AuthTestCollectionTeam1, collID)
}
}
})
t.Run("CollectionFilterIntegration", func(t *testing.T) {
// Test that collection-level filtering works in the search flow
t.Run("TeamMemberSearchOnlyFindsTeamData", func(t *testing.T) {
// UserA from Team1 searches - should ONLY find Team1 data
ctx := createAuthContext(TestUserA, TestTeam1, true, false)
// Filter collections - should only allow Team1 collection
// Filter collections first
allCollections := []string{AuthTestCollectionTeam1, AuthTestCollectionTeam2}
allowed := filterKBCollectionsByAuth(ctx, allCollections)
allowed := assistant.FilterKBCollectionsByAuth(ctx, allCollections)
// Should only allow Team1
assert.Contains(t, allowed, AuthTestCollectionTeam1)
assert.NotContains(t, allowed, AuthTestCollectionTeam2)
assert.Len(t, allowed, 1, "Should only have 1 allowed collection")
// Execute search on allowed collections only
cfg := &searchTypes.Config{
KB: &searchTypes.KBConfig{
Collections: allowed,
Threshold: 0.3,
},
// Search on allowed collections
result := executeKBSearchOnCollections(t, allowed, "quantum physics deep learning")
assert.Greater(t, len(result.Items), 0, "Should find Team1 documents")
// Verify ALL results are from Team1 collection only
for _, item := range result.Items {
assert.Equal(t, AuthTestCollectionTeam1, item.Collection,
"All results should be from Team1 collection, got: %s", item.Collection)
}
searcher := search.New(cfg, nil)
t.Logf(" ✓ Team1 member found %d items, all from Team1 collection", len(result.Items))
})
req := &searchTypes.Request{
Type: searchTypes.SearchTypeKB,
Query: "quantum physics",
Collections: allowed,
Threshold: 0.3,
Limit: 10,
Source: searchTypes.SourceAuto,
t.Run("TeamMemberCannotAccessOtherTeamData", func(t *testing.T) {
// UserA from Team1 tries to access Team2 - should be blocked
ctx := createAuthContext(TestUserA, TestTeam1, true, false)
// Try to filter Team2 collection
collections := []string{AuthTestCollectionTeam2}
allowed := assistant.FilterKBCollectionsByAuth(ctx, collections)
// Should be empty - no access
assert.Empty(t, allowed, "Team1 member should NOT have access to Team2 collection")
t.Log(" ✓ Team1 member correctly blocked from Team2 collection")
})
t.Run("OwnerSearchOnlyFindsOwnData", func(t *testing.T) {
// UserA with OwnerOnly - should only find collections they created
ctx := createAuthContext(TestUserA, "", false, true)
// Filter all collections
allCollections := []string{AuthTestCollectionTeam1, AuthTestCollectionTeam2, AuthTestCollectionPublic}
allowed := assistant.FilterKBCollectionsByAuth(ctx, allCollections)
// UserA created Team1 and Public, not Team2
assert.Contains(t, allowed, AuthTestCollectionTeam1, "Owner should access Team1 (created by UserA)")
assert.Contains(t, allowed, AuthTestCollectionPublic, "Owner should access Public (created by UserA)")
assert.NotContains(t, allowed, AuthTestCollectionTeam2, "Owner should NOT access Team2 (created by UserB)")
// Search and verify results
result := executeKBSearchOnCollections(t, allowed, "quantum artificial intelligence")
assert.Greater(t, len(result.Items), 0, "Should find owner's documents")
// Verify NO results from Team2
for _, item := range result.Items {
assert.NotEqual(t, AuthTestCollectionTeam2, item.Collection,
"Should NOT have results from Team2, got: %s", item.Collection)
}
t.Logf(" ✓ Owner found %d items, none from Team2", len(result.Items))
})
result, err := searcher.Search(nil, req)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.Greater(t, len(result.Items), 0, "Should find items in allowed collection")
t.Logf(" Found %d items in filtered collections", len(result.Items))
t.Run("NoConstraintsSearchFindsAllData", func(t *testing.T) {
// User with no constraints - should find all data
ctx := createAuthContext(TestUserA, TestTeam1, false, false)
// Filter all collections
allCollections := []string{AuthTestCollectionTeam1, AuthTestCollectionTeam2, AuthTestCollectionPublic}
allowed := assistant.FilterKBCollectionsByAuth(ctx, allCollections)
// Should have access to all
assert.Len(t, allowed, 3, "No constraints should allow all collections")
// Search and verify results from multiple collections
result := executeKBSearchOnCollections(t, allowed, "quantum deep learning artificial")
// Should find results from multiple collections
collectionsFound := make(map[string]bool)
for _, item := range result.Items {
collectionsFound[item.Collection] = true
}
assert.Greater(t, len(collectionsFound), 1, "Should find results from multiple collections")
t.Logf(" ✓ No constraints: found %d items from %d collections", len(result.Items), len(collectionsFound))
})
t.Run("SearchResultsMatchCollectionFilter", func(t *testing.T) {
// Verify that search results ONLY come from allowed collections
ctx := createAuthContext(TestUserB, TestTeam2, true, false)
// UserB from Team2 - should only access Team2
allCollections := []string{AuthTestCollectionTeam1, AuthTestCollectionTeam2, AuthTestCollectionPublic}
allowed := assistant.FilterKBCollectionsByAuth(ctx, allCollections)
assert.Contains(t, allowed, AuthTestCollectionTeam2, "Team2 member should access Team2")
assert.NotContains(t, allowed, AuthTestCollectionTeam1, "Team2 member should NOT access Team1")
// Search
result := executeKBSearchOnCollections(t, allowed, "deep learning computer vision")
// Verify results
if len(result.Items) > 0 {
for _, item := range result.Items {
// Results should only be from allowed collections
assert.Contains(t, allowed, item.Collection,
"Result from %s should be in allowed list %v", item.Collection, allowed)
}
t.Logf(" ✓ Team2 member found %d items, all from allowed collections", len(result.Items))
} else {
t.Log(" ✓ Team2 member found 0 items (collection may be empty)")
}
})
}
@ -449,9 +584,17 @@ func sanitizeForID(s string) string {
}
func executeKBSearch(t *testing.T, collectionID, query string, metadata map[string]interface{}) *searchTypes.Result {
return executeKBSearchOnCollections(t, []string{collectionID}, query)
}
func executeKBSearchOnCollections(t *testing.T, collections []string, query string) *searchTypes.Result {
if len(collections) == 0 {
return &searchTypes.Result{Items: []*searchTypes.ResultItem{}}
}
cfg := &searchTypes.Config{
KB: &searchTypes.KBConfig{
Collections: []string{collectionID},
Collections: collections,
Threshold: 0.3,
},
}
@ -460,11 +603,10 @@ func executeKBSearch(t *testing.T, collectionID, query string, metadata map[stri
req := &searchTypes.Request{
Type: searchTypes.SearchTypeKB,
Query: query,
Collections: []string{collectionID},
Collections: collections,
Threshold: 0.3,
Limit: 10,
Limit: 20,
Source: searchTypes.SourceAuto,
Metadata: metadata,
}
result, err := searcher.Search(nil, req)

View file

@ -8,10 +8,10 @@ import (
oauthtypes "github.com/yaoapp/yao/openapi/oauth/types"
)
// filterKBCollectionsByAuth filters collections based on user authorization.
// FilterKBCollectionsByAuth filters collections based on user authorization.
// Returns only collections that the user has permission to access.
// Permission is determined by Collection's metadata (public, share, __yao_team_id, __yao_created_by).
func filterKBCollectionsByAuth(ctx *agentContext.Context, collections []string) []string {
func FilterKBCollectionsByAuth(ctx *agentContext.Context, collections []string) []string {
if ctx == nil || ctx.Authorized == nil {
return collections // No auth context, return all
}