Implement authorized script execution and enhance context handling
- Added ExecuteWithAuthorized method to the Script struct, allowing scripts to be executed with user authorization information. - Updated existing Execute method to utilize ExecuteWithAuthorized for improved authorization handling. - Enhanced script execution tests to verify behavior with and without authorized information, ensuring accurate context management. - Implemented authorized information conversion to a map for easier integration with runtime environments. - Refactored context handling in script execution to ensure authorized data is correctly passed and utilized.
This commit is contained in:
parent
704c0331b1
commit
5f803ef4d4
10 changed files with 1339 additions and 5 deletions
|
|
@ -23,6 +23,11 @@ func (s *Script) Execute(ctx *context.Context, method string, args ...interface{
|
||||||
}
|
}
|
||||||
defer scriptCtx.Close()
|
defer scriptCtx.Close()
|
||||||
|
|
||||||
|
// Set authorized information if available
|
||||||
|
if ctx.Authorized != nil {
|
||||||
|
scriptCtx.WithAuthorized(ctx.Authorized.AuthorizedToMap())
|
||||||
|
}
|
||||||
|
|
||||||
// The first argument is the context
|
// The first argument is the context
|
||||||
args = append([]interface{}{ctx}, args...)
|
args = append([]interface{}{ctx}, args...)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -20,6 +20,11 @@ var scriptsMutex sync.Mutex
|
||||||
|
|
||||||
// Execute execute the script
|
// Execute execute the script
|
||||||
func (s *Script) Execute(ctx context.Context, method string, args ...interface{}) (interface{}, error) {
|
func (s *Script) Execute(ctx context.Context, method string, args ...interface{}) (interface{}, error) {
|
||||||
|
return s.ExecuteWithAuthorized(ctx, method, nil, args...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExecuteWithAuthorized execute the script with authorized information
|
||||||
|
func (s *Script) ExecuteWithAuthorized(ctx context.Context, method string, authorized map[string]interface{}, args ...interface{}) (interface{}, error) {
|
||||||
if s == nil || s.Script == nil {
|
if s == nil || s.Script == nil {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
@ -30,6 +35,11 @@ func (s *Script) Execute(ctx context.Context, method string, args ...interface{}
|
||||||
}
|
}
|
||||||
defer scriptCtx.Close()
|
defer scriptCtx.Close()
|
||||||
|
|
||||||
|
// Set authorized information if available
|
||||||
|
if authorized != nil {
|
||||||
|
scriptCtx.WithAuthorized(authorized)
|
||||||
|
}
|
||||||
|
|
||||||
// Call the method with provided arguments as-is
|
// Call the method with provided arguments as-is
|
||||||
result, err := scriptCtx.CallWith(ctx, method, args...)
|
result, err := scriptCtx.CallWith(ctx, method, args...)
|
||||||
|
|
||||||
|
|
@ -359,8 +369,14 @@ func makeScriptHandler(script *Script) process.Handler {
|
||||||
// Get arguments from process
|
// Get arguments from process
|
||||||
args := p.Args
|
args := p.Args
|
||||||
|
|
||||||
// Execute the script
|
// Convert authorized info to map if available
|
||||||
result, err := script.Execute(p.Context, method, args...)
|
var authorized map[string]interface{}
|
||||||
|
if p.Authorized != nil {
|
||||||
|
authorized = p.Authorized.AuthorizedToMap()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute the script with authorized information
|
||||||
|
result, err := script.ExecuteWithAuthorized(p.Context, method, authorized, args...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
exception.New(err.Error(), 500).Throw()
|
exception.New(err.Error(), 500).Throw()
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,10 +1,12 @@
|
||||||
package assistant
|
package assistant
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/yaoapp/gou/process"
|
||||||
"github.com/yaoapp/yao/config"
|
"github.com/yaoapp/yao/config"
|
||||||
"github.com/yaoapp/yao/test"
|
"github.com/yaoapp/yao/test"
|
||||||
)
|
)
|
||||||
|
|
@ -155,3 +157,151 @@ func TestGenerateScriptID(t *testing.T) {
|
||||||
// TestLoadScriptsThreadSafety tests concurrent script loading
|
// TestLoadScriptsThreadSafety tests concurrent script loading
|
||||||
// Note: This test is commented out due to path format differences
|
// Note: This test is commented out due to path format differences
|
||||||
// Thread safety is ensured by the scriptsMutex in LoadScripts function
|
// Thread safety is ensured by the scriptsMutex in LoadScripts function
|
||||||
|
|
||||||
|
func TestExecuteWithAuthorized(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
|
||||||
|
t.Run("ExecuteWithAuthorizedInfo", func(t *testing.T) {
|
||||||
|
// Create a script that returns the authorized info from __yao_data
|
||||||
|
scriptSource := `
|
||||||
|
function GetAuth() {
|
||||||
|
if (typeof __yao_data !== 'undefined' && __yao_data.AUTHORIZED) {
|
||||||
|
return __yao_data.AUTHORIZED;
|
||||||
|
}
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
`
|
||||||
|
|
||||||
|
data := map[string]interface{}{
|
||||||
|
"scripts": map[string]interface{}{
|
||||||
|
"auth_test": scriptSource,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, scripts, err := LoadScriptsFromData(data, "test.authorized")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, scripts)
|
||||||
|
require.Contains(t, scripts, "auth_test")
|
||||||
|
|
||||||
|
script := scripts["auth_test"]
|
||||||
|
|
||||||
|
// Create authorized info
|
||||||
|
authorized := map[string]interface{}{
|
||||||
|
"user_id": "user123",
|
||||||
|
"team_id": "team456",
|
||||||
|
"scope": "read write",
|
||||||
|
"constraints": map[string]interface{}{
|
||||||
|
"team_only": true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute with authorized info
|
||||||
|
ctx := context.Background()
|
||||||
|
result, err := script.ExecuteWithAuthorized(ctx, "GetAuth", authorized)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
// Verify the authorized info was passed correctly
|
||||||
|
resultMap, ok := result.(map[string]interface{})
|
||||||
|
require.True(t, ok, "Result should be a map")
|
||||||
|
|
||||||
|
assert.Equal(t, "user123", resultMap["user_id"])
|
||||||
|
assert.Equal(t, "team456", resultMap["team_id"])
|
||||||
|
assert.Equal(t, "read write", resultMap["scope"])
|
||||||
|
|
||||||
|
constraints, ok := resultMap["constraints"].(map[string]interface{})
|
||||||
|
require.True(t, ok, "Constraints should be a map")
|
||||||
|
assert.Equal(t, true, constraints["team_only"])
|
||||||
|
|
||||||
|
t.Logf("✓ Authorized info passed correctly to script")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("ExecuteWithoutAuthorizedInfo", func(t *testing.T) {
|
||||||
|
// Create a script that checks for authorized info
|
||||||
|
scriptSource := `
|
||||||
|
function CheckAuth() {
|
||||||
|
if (typeof __yao_data !== 'undefined' && __yao_data.AUTHORIZED) {
|
||||||
|
return { hasAuth: true, data: __yao_data.AUTHORIZED };
|
||||||
|
}
|
||||||
|
return { hasAuth: false };
|
||||||
|
}
|
||||||
|
`
|
||||||
|
|
||||||
|
data := map[string]interface{}{
|
||||||
|
"scripts": map[string]interface{}{
|
||||||
|
"no_auth_test": scriptSource,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, scripts, err := LoadScriptsFromData(data, "test.noauth")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, scripts)
|
||||||
|
require.Contains(t, scripts, "no_auth_test")
|
||||||
|
|
||||||
|
script := scripts["no_auth_test"]
|
||||||
|
|
||||||
|
// Execute without authorized info
|
||||||
|
ctx := context.Background()
|
||||||
|
result, err := script.Execute(ctx, "CheckAuth")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
resultMap, ok := result.(map[string]interface{})
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Equal(t, false, resultMap["hasAuth"])
|
||||||
|
|
||||||
|
t.Logf("✓ Script executed correctly without authorized info")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("MakeScriptHandlerWithAuthorized", func(t *testing.T) {
|
||||||
|
// Create a script that returns authorized user_id
|
||||||
|
scriptSource := `
|
||||||
|
function GetUserID() {
|
||||||
|
if (typeof __yao_data !== 'undefined' && __yao_data.AUTHORIZED) {
|
||||||
|
return __yao_data.AUTHORIZED.user_id || null;
|
||||||
|
}
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
`
|
||||||
|
|
||||||
|
data := map[string]interface{}{
|
||||||
|
"scripts": map[string]interface{}{
|
||||||
|
"handler_test": scriptSource,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, scripts, err := LoadScriptsFromData(data, "test.handler")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, scripts)
|
||||||
|
require.Contains(t, scripts, "handler_test")
|
||||||
|
|
||||||
|
script := scripts["handler_test"]
|
||||||
|
|
||||||
|
// Create a process handler
|
||||||
|
handler := makeScriptHandler(script)
|
||||||
|
require.NotNil(t, handler)
|
||||||
|
|
||||||
|
// Create a mock process with authorized info
|
||||||
|
ctx := context.Background()
|
||||||
|
p := &process.Process{
|
||||||
|
Method: "GetUserID",
|
||||||
|
Args: []interface{}{},
|
||||||
|
Context: ctx,
|
||||||
|
Authorized: &process.AuthorizedInfo{
|
||||||
|
UserID: "user999",
|
||||||
|
TeamID: "team888",
|
||||||
|
Scope: "admin",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute the handler
|
||||||
|
result := handler(p)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
// Verify the result
|
||||||
|
assert.Equal(t, "user999", result)
|
||||||
|
|
||||||
|
t.Logf("✓ Process handler correctly passed authorized info")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
|
||||||
373
openapi/kb/collection_process.go
Normal file
373
openapi/kb/collection_process.go
Normal file
|
|
@ -0,0 +1,373 @@
|
||||||
|
package kb
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
|
||||||
|
"github.com/yaoapp/gou/model"
|
||||||
|
"github.com/yaoapp/gou/process"
|
||||||
|
"github.com/yaoapp/kun/exception"
|
||||||
|
"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"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ProcessCreateCollection creates a new collection via Yao process
|
||||||
|
// Process: kb.collection.Create
|
||||||
|
//
|
||||||
|
// Args[0]: params (map) - Collection creation parameters
|
||||||
|
//
|
||||||
|
// {
|
||||||
|
// "id": "collection_id",
|
||||||
|
// "metadata": {
|
||||||
|
// "name": "Collection Name",
|
||||||
|
// "description": "Description"
|
||||||
|
// },
|
||||||
|
// "embedding_provider_id": "__yao.openai",
|
||||||
|
// "embedding_option_id": "text-embedding-3-small",
|
||||||
|
// "locale": "en",
|
||||||
|
// "config": {
|
||||||
|
// "distance": "cosine",
|
||||||
|
// "index_type": "hnsw",
|
||||||
|
// "m": 16,
|
||||||
|
// "ef_construction": 200,
|
||||||
|
// "ef_search": 64
|
||||||
|
// }
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// Returns: map with collection_id and message
|
||||||
|
func ProcessCreateCollection(process *process.Process) interface{} {
|
||||||
|
process.ValidateArgNums(1)
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
exception.New("Knowledge base not initialized", 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get authorized info from process
|
||||||
|
authInfo := authorized.ProcessAuthInfo(process)
|
||||||
|
|
||||||
|
// Parse parameters using JSON for type safety
|
||||||
|
paramsJSON, err := json.Marshal(process.Args[0])
|
||||||
|
if err != nil {
|
||||||
|
exception.New("Failed to encode parameters: "+err.Error(), 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
var params kbapi.CreateCollectionParams
|
||||||
|
if err := json.Unmarshal(paramsJSON, ¶ms); err != nil {
|
||||||
|
exception.New("Failed to decode parameters: "+err.Error(), 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply auth scope from authorized info
|
||||||
|
if authInfo != nil {
|
||||||
|
authScope := authInfo.WithCreateScope(maps.MapStrAny{})
|
||||||
|
params.AuthScope = authScope
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call API
|
||||||
|
result, err := kb.API.CreateCollection(process.Context, ¶ms)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to create collection: %v", err)
|
||||||
|
exception.New(err.Error(), 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
return maps.MapStrAny{
|
||||||
|
"collection_id": result.CollectionID,
|
||||||
|
"message": result.Message,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessRemoveCollection removes a collection via Yao process
|
||||||
|
// Process: kb.collection.Remove
|
||||||
|
//
|
||||||
|
// Args[0]: collection_id (string) - Collection ID to remove
|
||||||
|
//
|
||||||
|
// Returns: map with collection_id, removed status, documents_removed count, and message
|
||||||
|
func ProcessRemoveCollection(process *process.Process) interface{} {
|
||||||
|
process.ValidateArgNums(1)
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
exception.New("Knowledge base not initialized", 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get authorized info from process
|
||||||
|
authInfo := authorized.ProcessAuthInfo(process)
|
||||||
|
|
||||||
|
collectionID := process.ArgsString(0)
|
||||||
|
if collectionID == "" {
|
||||||
|
exception.New("Collection ID is required", 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check remove permission
|
||||||
|
hasPermission, err := checkCollectionPermission(authInfo, collectionID)
|
||||||
|
if err != nil {
|
||||||
|
exception.New(err.Error(), 403).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hasPermission {
|
||||||
|
exception.New("Forbidden: No permission to remove collection", 403).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := kb.API.RemoveCollection(process.Context, collectionID)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to remove collection: %v", err)
|
||||||
|
exception.New(err.Error(), 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
return maps.MapStrAny{
|
||||||
|
"collection_id": result.CollectionID,
|
||||||
|
"removed": result.Removed,
|
||||||
|
"documents_removed": result.DocumentsRemoved,
|
||||||
|
"message": result.Message,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessGetCollection retrieves a collection by ID via Yao process
|
||||||
|
// Process: kb.collection.Get
|
||||||
|
//
|
||||||
|
// Args[0]: collection_id (string) - Collection ID to retrieve
|
||||||
|
//
|
||||||
|
// Returns: map containing collection details
|
||||||
|
func ProcessGetCollection(process *process.Process) interface{} {
|
||||||
|
process.ValidateArgNums(1)
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
exception.New("Knowledge base not initialized", 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
collectionID := process.ArgsString(0)
|
||||||
|
if collectionID == "" {
|
||||||
|
exception.New("Collection ID is required", 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
collection, err := kb.API.GetCollection(process.Context, collectionID)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to get collection: %v", err)
|
||||||
|
exception.New(err.Error(), 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
return collection
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessCollectionExists checks if a collection exists via Yao process
|
||||||
|
// Process: kb.collection.Exists
|
||||||
|
//
|
||||||
|
// Args[0]: collection_id (string) - Collection ID to check
|
||||||
|
//
|
||||||
|
// Returns: map with collection_id and exists status
|
||||||
|
func ProcessCollectionExists(process *process.Process) interface{} {
|
||||||
|
process.ValidateArgNums(1)
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
exception.New("Knowledge base not initialized", 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
collectionID := process.ArgsString(0)
|
||||||
|
if collectionID == "" {
|
||||||
|
exception.New("Collection ID is required", 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := kb.API.CollectionExists(process.Context, collectionID)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to check collection existence: %v", err)
|
||||||
|
exception.New(err.Error(), 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
return maps.MapStrAny{
|
||||||
|
"collection_id": result.CollectionID,
|
||||||
|
"exists": result.Exists,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessListCollections lists collections with pagination via Yao process
|
||||||
|
// Process: kb.collection.List
|
||||||
|
//
|
||||||
|
// Args[0]: filter (map) - Optional filter parameters
|
||||||
|
//
|
||||||
|
// {
|
||||||
|
// "page": 1,
|
||||||
|
// "pagesize": 20,
|
||||||
|
// "keywords": "search term",
|
||||||
|
// "status": ["active"],
|
||||||
|
// "embedding_provider_id": "__yao.openai",
|
||||||
|
// "system": false,
|
||||||
|
// "select": ["id", "name", "status"],
|
||||||
|
// "sort": [{"column": "created_at", "option": "desc"}]
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// Returns: map with data array and pagination info
|
||||||
|
func ProcessListCollections(process *process.Process) interface{} {
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
exception.New("Knowledge base not initialized", 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get authorized info from process
|
||||||
|
authInfo := authorized.ProcessAuthInfo(process)
|
||||||
|
|
||||||
|
// Default filter
|
||||||
|
filter := &kbapi.ListCollectionsFilter{
|
||||||
|
Page: kbapi.DefaultPage,
|
||||||
|
PageSize: kbapi.DefaultPageSize,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse filter parameters using JSON (optional)
|
||||||
|
if process.NumOfArgs() > 0 {
|
||||||
|
filterJSON, err := json.Marshal(process.Args[0])
|
||||||
|
if err != nil {
|
||||||
|
exception.New("Failed to encode filter: "+err.Error(), 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(filterJSON, filter); err != nil {
|
||||||
|
exception.New("Failed to decode filter: "+err.Error(), 400).Throw()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply auth filters from authorized info
|
||||||
|
if authInfo != nil {
|
||||||
|
filter.AuthFilters = processAuthFilter(authInfo)
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := kb.API.ListCollections(process.Context, filter)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to list collections: %v", err)
|
||||||
|
exception.New(err.Error(), 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
return maps.MapStrAny{
|
||||||
|
"data": result.Data,
|
||||||
|
"next": result.Next,
|
||||||
|
"prev": result.Prev,
|
||||||
|
"page": result.Page,
|
||||||
|
"pagesize": result.PageSize,
|
||||||
|
"total": result.Total,
|
||||||
|
"pagecnt": result.PageCnt,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessUpdateCollectionMetadata updates collection metadata via Yao process
|
||||||
|
// Process: kb.collection.UpdateMetadata
|
||||||
|
//
|
||||||
|
// Args[0]: collection_id (string) - Collection ID
|
||||||
|
// Args[1]: params (map) - Update parameters
|
||||||
|
//
|
||||||
|
// {
|
||||||
|
// "metadata": {
|
||||||
|
// "name": "New Name",
|
||||||
|
// "description": "New Description"
|
||||||
|
// }
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// Returns: map with collection_id and message
|
||||||
|
func ProcessUpdateCollectionMetadata(process *process.Process) interface{} {
|
||||||
|
process.ValidateArgNums(2)
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
exception.New("Knowledge base not initialized", 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get authorized info from process
|
||||||
|
authInfo := authorized.ProcessAuthInfo(process)
|
||||||
|
|
||||||
|
collectionID := process.ArgsString(0)
|
||||||
|
if collectionID == "" {
|
||||||
|
exception.New("Collection ID is required", 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check update permission
|
||||||
|
hasPermission, err := checkCollectionPermission(authInfo, collectionID)
|
||||||
|
if err != nil {
|
||||||
|
exception.New(err.Error(), 403).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hasPermission {
|
||||||
|
exception.New("Forbidden: No permission to update collection", 403).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse parameters using JSON
|
||||||
|
paramsJSON, err := json.Marshal(process.Args[1])
|
||||||
|
if err != nil {
|
||||||
|
exception.New("Failed to encode parameters: "+err.Error(), 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
var params kbapi.UpdateMetadataParams
|
||||||
|
if err := json.Unmarshal(paramsJSON, ¶ms); err != nil {
|
||||||
|
exception.New("Failed to decode parameters: "+err.Error(), 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(params.Metadata) == 0 {
|
||||||
|
exception.New("Metadata is required and cannot be empty", 400).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply auth scope from authorized info
|
||||||
|
if authInfo != nil {
|
||||||
|
authScope := authInfo.WithUpdateScope(maps.MapStrAny{})
|
||||||
|
params.AuthScope = authScope
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := kb.API.UpdateCollectionMetadata(process.Context, collectionID, ¶ms)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to update collection metadata: %v", err)
|
||||||
|
exception.New(err.Error(), 500).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
return maps.MapStrAny{
|
||||||
|
"collection_id": result.CollectionID,
|
||||||
|
"message": result.Message,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper functions for Process handlers
|
||||||
|
|
||||||
|
// processAuthFilter applies permission-based filtering to query wheres for process handlers
|
||||||
|
// This function builds where clauses based on the user's authorization constraints
|
||||||
|
func processAuthFilter(authInfo *oauthtypes.AuthorizedInfo) []model.QueryWhere {
|
||||||
|
if authInfo == nil {
|
||||||
|
return []model.QueryWhere{}
|
||||||
|
}
|
||||||
|
|
||||||
|
var wheres []model.QueryWhere
|
||||||
|
scope := authInfo.AccessScope()
|
||||||
|
|
||||||
|
// Team only - User can access:
|
||||||
|
// 1. Public records (public = true)
|
||||||
|
// 2. Records in their team where:
|
||||||
|
// - They created the record (__yao_created_by matches)
|
||||||
|
// - OR the record is shared with team (share = "team")
|
||||||
|
if authInfo.Constraints.TeamOnly && authInfo.TeamID != "" && authInfo.UserID != "" {
|
||||||
|
wheres = append(wheres, model.QueryWhere{
|
||||||
|
Wheres: []model.QueryWhere{
|
||||||
|
{Column: "public", Value: true, Method: "orwhere"},
|
||||||
|
{Wheres: []model.QueryWhere{
|
||||||
|
{Column: "__yao_team_id", Value: scope.TeamID},
|
||||||
|
{Wheres: []model.QueryWhere{
|
||||||
|
{Column: "__yao_created_by", Value: scope.CreatedBy},
|
||||||
|
{Column: "share", Value: "team", Method: "orwhere"},
|
||||||
|
}},
|
||||||
|
}, Method: "orwhere"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return wheres
|
||||||
|
}
|
||||||
|
|
||||||
|
// Owner only - User can access:
|
||||||
|
// 1. Public records (public = true)
|
||||||
|
// 2. Records they created where:
|
||||||
|
// - __yao_team_id is null (not team records)
|
||||||
|
// - __yao_created_by matches their user ID
|
||||||
|
if authInfo.Constraints.OwnerOnly && authInfo.UserID != "" {
|
||||||
|
wheres = append(wheres, model.QueryWhere{
|
||||||
|
Wheres: []model.QueryWhere{
|
||||||
|
{Column: "public", Value: true, Method: "orwhere"},
|
||||||
|
{Wheres: []model.QueryWhere{
|
||||||
|
{Column: "__yao_team_id", OP: "null"},
|
||||||
|
{Column: "__yao_created_by", Value: scope.CreatedBy},
|
||||||
|
}, Method: "orwhere"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return wheres
|
||||||
|
}
|
||||||
|
|
||||||
|
return wheres
|
||||||
|
}
|
||||||
|
|
@ -11,6 +11,15 @@ import (
|
||||||
func init() {
|
func init() {
|
||||||
// Register kb process handlers
|
// Register kb process handlers
|
||||||
process.RegisterGroup("kb", map[string]process.Handler{
|
process.RegisterGroup("kb", map[string]process.Handler{
|
||||||
|
// Collection processes
|
||||||
|
"collection.create": ProcessCreateCollection,
|
||||||
|
"collection.remove": ProcessRemoveCollection,
|
||||||
|
"collection.get": ProcessGetCollection,
|
||||||
|
"collection.exists": ProcessCollectionExists,
|
||||||
|
"collection.list": ProcessListCollections,
|
||||||
|
"collection.updatemetadata": ProcessUpdateCollectionMetadata,
|
||||||
|
|
||||||
|
// Document processes
|
||||||
"documents.addfile": ProcessAddFile,
|
"documents.addfile": ProcessAddFile,
|
||||||
"documents.addtext": ProcessAddText,
|
"documents.addtext": ProcessAddText,
|
||||||
"documents.addurl": ProcessAddURL,
|
"documents.addurl": ProcessAddURL,
|
||||||
|
|
|
||||||
|
|
@ -8,9 +8,37 @@ import (
|
||||||
|
|
||||||
// ProcessAuthInfo extracts authorized information from the process
|
// ProcessAuthInfo extracts authorized information from the process
|
||||||
func ProcessAuthInfo(p *process.Process) *types.AuthorizedInfo {
|
func ProcessAuthInfo(p *process.Process) *types.AuthorizedInfo {
|
||||||
// TODO: Implement this function
|
if p == nil {
|
||||||
// Get authorized information from the process context
|
return nil
|
||||||
info := &types.AuthorizedInfo{}
|
}
|
||||||
|
|
||||||
|
// Get authorized info from process
|
||||||
|
processAuth := p.GetAuthorized()
|
||||||
|
if processAuth == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert process.AuthorizedInfo to types.AuthorizedInfo
|
||||||
|
info := &types.AuthorizedInfo{
|
||||||
|
Subject: processAuth.Subject,
|
||||||
|
ClientID: processAuth.ClientID,
|
||||||
|
UserID: processAuth.UserID,
|
||||||
|
Scope: processAuth.Scope,
|
||||||
|
TeamID: processAuth.TeamID,
|
||||||
|
TenantID: processAuth.TenantID,
|
||||||
|
SessionID: processAuth.SessionID,
|
||||||
|
RememberMe: processAuth.RememberMe,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert constraints
|
||||||
|
info.Constraints = types.DataConstraints{
|
||||||
|
OwnerOnly: processAuth.Constraints.OwnerOnly,
|
||||||
|
CreatorOnly: processAuth.Constraints.CreatorOnly,
|
||||||
|
EditorOnly: processAuth.Constraints.EditorOnly,
|
||||||
|
TeamOnly: processAuth.Constraints.TeamOnly,
|
||||||
|
Extra: processAuth.Constraints.Extra,
|
||||||
|
}
|
||||||
|
|
||||||
return info
|
return info
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
88
openapi/oauth/authorized/utils_test.go
Normal file
88
openapi/oauth/authorized/utils_test.go
Normal file
|
|
@ -0,0 +1,88 @@
|
||||||
|
package authorized
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/yaoapp/gou/process"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestProcessAuthInfo(t *testing.T) {
|
||||||
|
t.Run("WithNilProcess", func(t *testing.T) {
|
||||||
|
result := ProcessAuthInfo(nil)
|
||||||
|
assert.Nil(t, result)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("WithProcessNoAuth", func(t *testing.T) {
|
||||||
|
p := &process.Process{}
|
||||||
|
result := ProcessAuthInfo(p)
|
||||||
|
// GetAuthorized returns empty struct instead of nil, so ProcessAuthInfo will return an empty AuthorizedInfo
|
||||||
|
require.NotNil(t, result)
|
||||||
|
assert.Empty(t, result.UserID)
|
||||||
|
assert.Empty(t, result.TeamID)
|
||||||
|
assert.Empty(t, result.Subject)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("WithProcessWithAuth", func(t *testing.T) {
|
||||||
|
p := &process.Process{
|
||||||
|
Authorized: &process.AuthorizedInfo{
|
||||||
|
Subject: "user123",
|
||||||
|
ClientID: "client456",
|
||||||
|
UserID: "u789",
|
||||||
|
Scope: "read write",
|
||||||
|
TeamID: "t123",
|
||||||
|
TenantID: "tenant456",
|
||||||
|
SessionID: "session789",
|
||||||
|
RememberMe: true,
|
||||||
|
Constraints: process.DataConstraints{
|
||||||
|
OwnerOnly: true,
|
||||||
|
CreatorOnly: false,
|
||||||
|
EditorOnly: false,
|
||||||
|
TeamOnly: true,
|
||||||
|
Extra: map[string]interface{}{
|
||||||
|
"department": "engineering",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ProcessAuthInfo(p)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
assert.Equal(t, "user123", result.Subject)
|
||||||
|
assert.Equal(t, "client456", result.ClientID)
|
||||||
|
assert.Equal(t, "u789", result.UserID)
|
||||||
|
assert.Equal(t, "read write", result.Scope)
|
||||||
|
assert.Equal(t, "t123", result.TeamID)
|
||||||
|
assert.Equal(t, "tenant456", result.TenantID)
|
||||||
|
assert.Equal(t, "session789", result.SessionID)
|
||||||
|
assert.True(t, result.RememberMe)
|
||||||
|
|
||||||
|
assert.True(t, result.Constraints.OwnerOnly)
|
||||||
|
assert.False(t, result.Constraints.CreatorOnly)
|
||||||
|
assert.False(t, result.Constraints.EditorOnly)
|
||||||
|
assert.True(t, result.Constraints.TeamOnly)
|
||||||
|
assert.Equal(t, "engineering", result.Constraints.Extra["department"])
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("WithPartialData", func(t *testing.T) {
|
||||||
|
p := &process.Process{
|
||||||
|
Authorized: &process.AuthorizedInfo{
|
||||||
|
UserID: "u123",
|
||||||
|
TeamID: "t456",
|
||||||
|
Constraints: process.DataConstraints{
|
||||||
|
TeamOnly: true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ProcessAuthInfo(p)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
assert.Equal(t, "u123", result.UserID)
|
||||||
|
assert.Equal(t, "t456", result.TeamID)
|
||||||
|
assert.True(t, result.Constraints.TeamOnly)
|
||||||
|
assert.False(t, result.Constraints.OwnerOnly)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
@ -462,3 +462,152 @@ func TestCopyScopesIntegration(t *testing.T) {
|
||||||
assert.Equal(t, "tenant789", updateResult["__yao_tenant_id"])
|
assert.Equal(t, "tenant789", updateResult["__yao_tenant_id"])
|
||||||
assert.Nil(t, updateResult["__yao_created_by"]) // Should not be copied
|
assert.Nil(t, updateResult["__yao_created_by"]) // Should not be copied
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAuthorizedToMap(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
auth *AuthorizedInfo
|
||||||
|
expected map[string]interface{}
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "Full AuthorizedInfo",
|
||||||
|
auth: &AuthorizedInfo{
|
||||||
|
Subject: "user123",
|
||||||
|
ClientID: "client456",
|
||||||
|
Scope: "read write",
|
||||||
|
SessionID: "session789",
|
||||||
|
UserID: "user123",
|
||||||
|
TeamID: "team456",
|
||||||
|
TenantID: "tenant789",
|
||||||
|
RememberMe: true,
|
||||||
|
Constraints: DataConstraints{
|
||||||
|
OwnerOnly: true,
|
||||||
|
CreatorOnly: false,
|
||||||
|
EditorOnly: false,
|
||||||
|
TeamOnly: true,
|
||||||
|
Extra: map[string]interface{}{
|
||||||
|
"department": "engineering",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
expected: map[string]interface{}{
|
||||||
|
"sub": "user123",
|
||||||
|
"client_id": "client456",
|
||||||
|
"scope": "read write",
|
||||||
|
"session_id": "session789",
|
||||||
|
"user_id": "user123",
|
||||||
|
"team_id": "team456",
|
||||||
|
"tenant_id": "tenant789",
|
||||||
|
"remember_me": true,
|
||||||
|
"constraints": map[string]interface{}{
|
||||||
|
"owner_only": true,
|
||||||
|
"team_only": true,
|
||||||
|
"extra": map[string]interface{}{
|
||||||
|
"department": "engineering",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Partial AuthorizedInfo",
|
||||||
|
auth: &AuthorizedInfo{
|
||||||
|
UserID: "user123",
|
||||||
|
TeamID: "team456",
|
||||||
|
},
|
||||||
|
expected: map[string]interface{}{
|
||||||
|
"user_id": "user123",
|
||||||
|
"team_id": "team456",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "AuthorizedInfo with only constraints",
|
||||||
|
auth: &AuthorizedInfo{
|
||||||
|
UserID: "user123",
|
||||||
|
Constraints: DataConstraints{
|
||||||
|
TeamOnly: true,
|
||||||
|
Extra: map[string]interface{}{
|
||||||
|
"region": "us-west",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
expected: map[string]interface{}{
|
||||||
|
"user_id": "user123",
|
||||||
|
"constraints": map[string]interface{}{
|
||||||
|
"team_only": true,
|
||||||
|
"extra": map[string]interface{}{
|
||||||
|
"region": "us-west",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Nil AuthorizedInfo",
|
||||||
|
auth: nil,
|
||||||
|
expected: nil,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Empty AuthorizedInfo",
|
||||||
|
auth: &AuthorizedInfo{},
|
||||||
|
expected: map[string]interface{}{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
result := tt.auth.AuthorizedToMap()
|
||||||
|
|
||||||
|
if tt.expected == nil {
|
||||||
|
assert.Nil(t, result)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
|
||||||
|
// Check all expected keys
|
||||||
|
for key, expectedValue := range tt.expected {
|
||||||
|
actualValue, ok := result[key]
|
||||||
|
if !ok {
|
||||||
|
t.Errorf("Key %s not found in result", key)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Special handling for nested maps (constraints)
|
||||||
|
if key == "constraints" {
|
||||||
|
expectedConstraints, _ := expectedValue.(map[string]interface{})
|
||||||
|
actualConstraints, ok := actualValue.(map[string]interface{})
|
||||||
|
assert.True(t, ok, "constraints should be map[string]interface{}")
|
||||||
|
|
||||||
|
for cKey, cExpectedValue := range expectedConstraints {
|
||||||
|
cActualValue, ok := actualConstraints[cKey]
|
||||||
|
assert.True(t, ok, "Constraint key %s should exist", cKey)
|
||||||
|
|
||||||
|
// Special handling for nested extra map
|
||||||
|
if cKey == "extra" {
|
||||||
|
expectedExtra, _ := cExpectedValue.(map[string]interface{})
|
||||||
|
actualExtra, ok := cActualValue.(map[string]interface{})
|
||||||
|
assert.True(t, ok, "extra should be map[string]interface{}")
|
||||||
|
|
||||||
|
for eKey, eExpectedValue := range expectedExtra {
|
||||||
|
eActualValue, ok := actualExtra[eKey]
|
||||||
|
assert.True(t, ok, "Extra key %s should exist", eKey)
|
||||||
|
assert.Equal(t, eExpectedValue, eActualValue)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
assert.Equal(t, cExpectedValue, cActualValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
assert.Equal(t, expectedValue, actualValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check no unexpected keys (except for empty maps)
|
||||||
|
if len(tt.expected) > 0 {
|
||||||
|
for key := range result {
|
||||||
|
_, ok := tt.expected[key]
|
||||||
|
assert.True(t, ok, "Unexpected key %s in result", key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -627,6 +627,64 @@ type AuthorizedInfo struct {
|
||||||
Constraints DataConstraints `json:"constraints,omitempty"`
|
Constraints DataConstraints `json:"constraints,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AuthorizedToMap converts AuthorizedInfo to map[string]interface{}
|
||||||
|
// This is useful for passing authorized information to runtime bridges (e.g., V8)
|
||||||
|
func (auth *AuthorizedInfo) AuthorizedToMap() map[string]interface{} {
|
||||||
|
if auth == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
result := make(map[string]interface{})
|
||||||
|
|
||||||
|
if auth.Subject != "" {
|
||||||
|
result["sub"] = auth.Subject
|
||||||
|
}
|
||||||
|
if auth.ClientID != "" {
|
||||||
|
result["client_id"] = auth.ClientID
|
||||||
|
}
|
||||||
|
if auth.Scope != "" {
|
||||||
|
result["scope"] = auth.Scope
|
||||||
|
}
|
||||||
|
if auth.SessionID != "" {
|
||||||
|
result["session_id"] = auth.SessionID
|
||||||
|
}
|
||||||
|
if auth.UserID != "" {
|
||||||
|
result["user_id"] = auth.UserID
|
||||||
|
}
|
||||||
|
if auth.TeamID != "" {
|
||||||
|
result["team_id"] = auth.TeamID
|
||||||
|
}
|
||||||
|
if auth.TenantID != "" {
|
||||||
|
result["tenant_id"] = auth.TenantID
|
||||||
|
}
|
||||||
|
if auth.RememberMe {
|
||||||
|
result["remember_me"] = auth.RememberMe
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add constraints if any are set
|
||||||
|
if auth.Constraints.OwnerOnly || auth.Constraints.CreatorOnly || auth.Constraints.EditorOnly || auth.Constraints.TeamOnly || len(auth.Constraints.Extra) > 0 {
|
||||||
|
constraints := make(map[string]interface{})
|
||||||
|
if auth.Constraints.OwnerOnly {
|
||||||
|
constraints["owner_only"] = true
|
||||||
|
}
|
||||||
|
if auth.Constraints.CreatorOnly {
|
||||||
|
constraints["creator_only"] = true
|
||||||
|
}
|
||||||
|
if auth.Constraints.EditorOnly {
|
||||||
|
constraints["editor_only"] = true
|
||||||
|
}
|
||||||
|
if auth.Constraints.TeamOnly {
|
||||||
|
constraints["team_only"] = true
|
||||||
|
}
|
||||||
|
if len(auth.Constraints.Extra) > 0 {
|
||||||
|
constraints["extra"] = auth.Constraints.Extra
|
||||||
|
}
|
||||||
|
result["constraints"] = constraints
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
// JWTClaims represents JWT-specific claims structure
|
// JWTClaims represents JWT-specific claims structure
|
||||||
type JWTClaims struct {
|
type JWTClaims struct {
|
||||||
jwt.StandardClaims
|
jwt.StandardClaims
|
||||||
|
|
|
||||||
458
openapi/tests/kb/collection_process_test.go
Normal file
458
openapi/tests/kb/collection_process_test.go
Normal file
|
|
@ -0,0 +1,458 @@
|
||||||
|
package openapi_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/yaoapp/gou/process"
|
||||||
|
"github.com/yaoapp/kun/maps"
|
||||||
|
"github.com/yaoapp/yao/kb"
|
||||||
|
"github.com/yaoapp/yao/openapi/tests/testutils"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestProcessCreateCollection tests the kb.collection.Create process
|
||||||
|
func TestProcessCreateCollection(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean()
|
||||||
|
|
||||||
|
// Ensure KB is initialized
|
||||||
|
if kb.API == nil {
|
||||||
|
t.Skip("Knowledge base not initialized - skipping test")
|
||||||
|
}
|
||||||
|
|
||||||
|
testCollectionID := fmt.Sprintf("test_process_create_%d", time.Now().UnixNano())
|
||||||
|
|
||||||
|
t.Run("CreateCollectionWithAuth", func(t *testing.T) {
|
||||||
|
// Register collection for cleanup
|
||||||
|
testutils.RegisterTestCollection(testCollectionID)
|
||||||
|
|
||||||
|
// Create process with authorized info
|
||||||
|
p := process.New("kb.collection.Create").
|
||||||
|
WithContext(context.Background()).
|
||||||
|
WithAuthorized(&process.AuthorizedInfo{
|
||||||
|
UserID: "test_user_123",
|
||||||
|
TeamID: "test_team_456",
|
||||||
|
Subject: "user@example.com",
|
||||||
|
ClientID: "test_client",
|
||||||
|
Scope: "openid profile",
|
||||||
|
Constraints: process.DataConstraints{
|
||||||
|
TeamOnly: true,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
// Prepare parameters - use map for Process API
|
||||||
|
params := map[string]interface{}{
|
||||||
|
"id": testCollectionID,
|
||||||
|
"metadata": map[string]interface{}{
|
||||||
|
"name": "Process Test Collection",
|
||||||
|
"description": "Created via Process API with auth",
|
||||||
|
},
|
||||||
|
"embedding_provider_id": "__yao.openai",
|
||||||
|
"embedding_option_id": "text-embedding-3-small",
|
||||||
|
"locale": "en",
|
||||||
|
"config": map[string]interface{}{
|
||||||
|
"index_type": "hnsw",
|
||||||
|
"distance": "cosine",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
p.Args = []interface{}{params}
|
||||||
|
|
||||||
|
// Execute process
|
||||||
|
result, err := p.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
// Verify result
|
||||||
|
resultMap, ok := result.(maps.MapStrAny)
|
||||||
|
require.True(t, ok, "Result should be a maps.MapStrAny")
|
||||||
|
assert.Equal(t, testCollectionID, resultMap["collection_id"])
|
||||||
|
assert.Contains(t, resultMap, "message")
|
||||||
|
|
||||||
|
t.Logf("✓ Successfully created collection via process: %s", testCollectionID)
|
||||||
|
|
||||||
|
// Verify auth scope was applied by checking the collection
|
||||||
|
collection, err := kb.API.GetCollection(context.Background(), testCollectionID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NotNil(t, collection)
|
||||||
|
|
||||||
|
// Check if auth fields were set
|
||||||
|
if createdBy, ok := collection["__yao_created_by"]; ok {
|
||||||
|
assert.Equal(t, "test_user_123", createdBy)
|
||||||
|
t.Logf("✓ Auth scope applied: __yao_created_by = %v", createdBy)
|
||||||
|
}
|
||||||
|
if teamID, ok := collection["__yao_team_id"]; ok {
|
||||||
|
assert.Equal(t, "test_team_456", teamID)
|
||||||
|
t.Logf("✓ Auth scope applied: __yao_team_id = %v", teamID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("CreateCollectionWithoutAuth", func(t *testing.T) {
|
||||||
|
testCollectionID2 := fmt.Sprintf("test_process_create_noauth_%d", time.Now().UnixNano())
|
||||||
|
testutils.RegisterTestCollection(testCollectionID2)
|
||||||
|
|
||||||
|
// Create process without authorized info
|
||||||
|
p := process.New("kb.collection.Create").
|
||||||
|
WithContext(context.Background())
|
||||||
|
|
||||||
|
params := map[string]interface{}{
|
||||||
|
"id": testCollectionID2,
|
||||||
|
"metadata": map[string]interface{}{
|
||||||
|
"name": "Process Test Collection No Auth",
|
||||||
|
"description": "Created via Process API without auth",
|
||||||
|
},
|
||||||
|
"embedding_provider_id": "__yao.openai",
|
||||||
|
"embedding_option_id": "text-embedding-3-small",
|
||||||
|
"locale": "en",
|
||||||
|
"config": map[string]interface{}{
|
||||||
|
"index_type": "hnsw",
|
||||||
|
"distance": "cosine",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
p.Args = []interface{}{params}
|
||||||
|
|
||||||
|
// Execute process
|
||||||
|
result, err := p.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
resultMap, ok := result.(maps.MapStrAny)
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Equal(t, testCollectionID2, resultMap["collection_id"])
|
||||||
|
|
||||||
|
t.Logf("✓ Successfully created collection without auth: %s", testCollectionID2)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("CreateCollectionInvalidParams", func(t *testing.T) {
|
||||||
|
// Create process with invalid parameters
|
||||||
|
p := process.New("kb.collection.Create").
|
||||||
|
WithContext(context.Background())
|
||||||
|
|
||||||
|
// Missing required fields
|
||||||
|
params := map[string]interface{}{
|
||||||
|
"metadata": map[string]interface{}{
|
||||||
|
"name": "Invalid Collection",
|
||||||
|
},
|
||||||
|
// Missing id, embedding_provider_id, etc.
|
||||||
|
}
|
||||||
|
|
||||||
|
p.Args = []interface{}{params}
|
||||||
|
|
||||||
|
// Execute should throw exception or return error
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
t.Logf("✓ Correctly rejected invalid parameters via panic: %v", r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
result, err := p.Exec()
|
||||||
|
if err != nil {
|
||||||
|
t.Logf("✓ Correctly rejected invalid parameters via error: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if result == nil {
|
||||||
|
t.Log("✓ Correctly rejected invalid parameters (nil result)")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Error("Should have thrown exception or returned error for invalid parameters")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestProcessListCollections tests the kb.collection.List process
|
||||||
|
func TestProcessListCollections(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean()
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
t.Skip("Knowledge base not initialized - skipping test")
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("ListCollectionsWithAuth", func(t *testing.T) {
|
||||||
|
// Test listing with auth filters
|
||||||
|
p := process.New("kb.collection.List").
|
||||||
|
WithContext(context.Background()).
|
||||||
|
WithAuthorized(&process.AuthorizedInfo{
|
||||||
|
UserID: "test_user_789",
|
||||||
|
TeamID: "test_team_789",
|
||||||
|
Constraints: process.DataConstraints{
|
||||||
|
TeamOnly: true,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
filter := map[string]interface{}{
|
||||||
|
"page": 1,
|
||||||
|
"pagesize": 20,
|
||||||
|
}
|
||||||
|
|
||||||
|
p.Args = []interface{}{filter}
|
||||||
|
|
||||||
|
// Execute process
|
||||||
|
result, err := p.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
resultMap, ok := result.(maps.MapStrAny)
|
||||||
|
require.True(t, ok, "Result should be a map")
|
||||||
|
|
||||||
|
assert.Contains(t, resultMap, "data")
|
||||||
|
assert.Contains(t, resultMap, "page")
|
||||||
|
assert.Contains(t, resultMap, "pagesize")
|
||||||
|
assert.Contains(t, resultMap, "total")
|
||||||
|
|
||||||
|
t.Logf("✓ Retrieved collections with auth filters")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("ListCollectionsNoFilter", func(t *testing.T) {
|
||||||
|
// Test listing without filter (should use defaults)
|
||||||
|
p := process.New("kb.collection.List").
|
||||||
|
WithContext(context.Background())
|
||||||
|
|
||||||
|
// No arguments - should use default filter
|
||||||
|
p.Args = []interface{}{}
|
||||||
|
|
||||||
|
result, err := p.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
resultMap, ok := result.(maps.MapStrAny)
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Contains(t, resultMap, "data")
|
||||||
|
assert.Equal(t, 1, resultMap["page"])
|
||||||
|
assert.Equal(t, 20, resultMap["pagesize"])
|
||||||
|
|
||||||
|
t.Logf("✓ Retrieved collections with default filter")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestProcessGetCollection tests the kb.collection.Get process
|
||||||
|
func TestProcessGetCollection(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean()
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
t.Skip("Knowledge base not initialized - skipping test")
|
||||||
|
}
|
||||||
|
|
||||||
|
testCollectionID := fmt.Sprintf("test_process_get_%d", time.Now().UnixNano())
|
||||||
|
testutils.RegisterTestCollection(testCollectionID)
|
||||||
|
|
||||||
|
t.Run("GetCollectionNotFound", func(t *testing.T) {
|
||||||
|
p := process.New("kb.collection.Get").
|
||||||
|
WithContext(context.Background())
|
||||||
|
|
||||||
|
p.Args = []interface{}{"nonexistent_collection_id"}
|
||||||
|
|
||||||
|
// Execute should throw exception or return error
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
t.Logf("✓ Correctly rejected nonexistent collection via panic: %v", r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
result, err := p.Exec()
|
||||||
|
if err != nil {
|
||||||
|
t.Logf("✓ Correctly rejected nonexistent collection via error: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if result == nil {
|
||||||
|
t.Log("✓ Correctly rejected nonexistent collection (nil result)")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Error("Should have thrown exception or returned error for nonexistent collection")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestProcessCollectionExists tests the kb.collection.Exists process
|
||||||
|
func TestProcessCollectionExists(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean()
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
t.Skip("Knowledge base not initialized - skipping test")
|
||||||
|
}
|
||||||
|
|
||||||
|
testCollectionID := fmt.Sprintf("test_process_exists_%d", time.Now().UnixNano())
|
||||||
|
testutils.RegisterTestCollection(testCollectionID)
|
||||||
|
|
||||||
|
t.Run("CollectionExistsBeforeCreation", func(t *testing.T) {
|
||||||
|
p := process.New("kb.collection.Exists").
|
||||||
|
WithContext(context.Background())
|
||||||
|
|
||||||
|
p.Args = []interface{}{testCollectionID}
|
||||||
|
|
||||||
|
result, err := p.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
resultMap, ok := result.(maps.MapStrAny)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
assert.Equal(t, testCollectionID, resultMap["collection_id"])
|
||||||
|
assert.Equal(t, false, resultMap["exists"])
|
||||||
|
|
||||||
|
t.Logf("✓ Correctly reported collection does not exist")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestProcessCollectionIntegration tests the full collection lifecycle via Process API
|
||||||
|
func TestProcessCollectionIntegration(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean()
|
||||||
|
|
||||||
|
if kb.API == nil {
|
||||||
|
t.Skip("Knowledge base not initialized - skipping test")
|
||||||
|
}
|
||||||
|
|
||||||
|
testCollectionID := fmt.Sprintf("test_process_integration_%d", time.Now().UnixNano())
|
||||||
|
testutils.RegisterTestCollection(testCollectionID)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
authInfo := &process.AuthorizedInfo{
|
||||||
|
UserID: "integration_user",
|
||||||
|
TeamID: "integration_team",
|
||||||
|
Subject: "integration@example.com",
|
||||||
|
ClientID: "integration_client",
|
||||||
|
Scope: "openid profile",
|
||||||
|
Constraints: process.DataConstraints{
|
||||||
|
TeamOnly: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("FullLifecycleViaProcess", func(t *testing.T) {
|
||||||
|
// Step 1: Check collection doesn't exist
|
||||||
|
p1 := process.New("kb.collection.Exists").WithContext(ctx)
|
||||||
|
p1.Args = []interface{}{testCollectionID}
|
||||||
|
result1, err := p1.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
existsResult := result1.(maps.MapStrAny)
|
||||||
|
assert.Equal(t, false, existsResult["exists"])
|
||||||
|
t.Logf("✓ Step 1: Confirmed collection doesn't exist")
|
||||||
|
|
||||||
|
// Step 2: Create collection
|
||||||
|
p2 := process.New("kb.collection.Create").WithContext(ctx).WithAuthorized(authInfo)
|
||||||
|
p2.Args = []interface{}{
|
||||||
|
map[string]interface{}{
|
||||||
|
"id": testCollectionID,
|
||||||
|
"metadata": map[string]interface{}{
|
||||||
|
"name": "Integration Test Collection",
|
||||||
|
"description": "Full lifecycle test",
|
||||||
|
},
|
||||||
|
"embedding_provider_id": "__yao.openai",
|
||||||
|
"embedding_option_id": "text-embedding-3-small",
|
||||||
|
"locale": "en",
|
||||||
|
"config": map[string]interface{}{
|
||||||
|
"index_type": "hnsw",
|
||||||
|
"distance": "cosine",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
result2, err := p2.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
createResult := result2.(maps.MapStrAny)
|
||||||
|
assert.Equal(t, testCollectionID, createResult["collection_id"])
|
||||||
|
t.Logf("✓ Step 2: Created collection")
|
||||||
|
|
||||||
|
// Step 3: Verify collection exists
|
||||||
|
p3 := process.New("kb.collection.Exists").WithContext(ctx)
|
||||||
|
p3.Args = []interface{}{testCollectionID}
|
||||||
|
result3, err := p3.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
existsResult2 := result3.(maps.MapStrAny)
|
||||||
|
assert.Equal(t, true, existsResult2["exists"])
|
||||||
|
t.Logf("✓ Step 3: Confirmed collection exists")
|
||||||
|
|
||||||
|
// Step 4: Get collection
|
||||||
|
p4 := process.New("kb.collection.Get").WithContext(ctx)
|
||||||
|
p4.Args = []interface{}{testCollectionID}
|
||||||
|
result4, err := p4.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
// GetCollection returns map[string]interface{}
|
||||||
|
var getResult map[string]interface{}
|
||||||
|
if mapStrAny, ok := result4.(maps.MapStrAny); ok {
|
||||||
|
getResult = mapStrAny
|
||||||
|
} else if m, ok := result4.(map[string]interface{}); ok {
|
||||||
|
getResult = m
|
||||||
|
} else {
|
||||||
|
t.Fatalf("Unexpected result type: %T", result4)
|
||||||
|
}
|
||||||
|
assert.Equal(t, "Integration Test Collection", getResult["name"])
|
||||||
|
t.Logf("✓ Step 4: Retrieved collection details")
|
||||||
|
|
||||||
|
// Step 5: Update metadata
|
||||||
|
p5 := process.New("kb.collection.UpdateMetadata").WithContext(ctx).WithAuthorized(authInfo)
|
||||||
|
p5.Args = []interface{}{
|
||||||
|
testCollectionID,
|
||||||
|
map[string]interface{}{
|
||||||
|
"metadata": map[string]interface{}{
|
||||||
|
"name": "Updated Integration Collection",
|
||||||
|
"description": "Updated via process",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
result5, err := p5.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
updateResult := result5.(maps.MapStrAny)
|
||||||
|
assert.Equal(t, testCollectionID, updateResult["collection_id"])
|
||||||
|
t.Logf("✓ Step 5: Updated collection metadata")
|
||||||
|
|
||||||
|
// Step 6: Verify update
|
||||||
|
p6 := process.New("kb.collection.Get").WithContext(ctx)
|
||||||
|
p6.Args = []interface{}{testCollectionID}
|
||||||
|
result6, err := p6.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
// GetCollection returns map[string]interface{}
|
||||||
|
var getResult2 map[string]interface{}
|
||||||
|
if mapStrAny, ok := result6.(maps.MapStrAny); ok {
|
||||||
|
getResult2 = mapStrAny
|
||||||
|
} else if m, ok := result6.(map[string]interface{}); ok {
|
||||||
|
getResult2 = m
|
||||||
|
} else {
|
||||||
|
t.Fatalf("Unexpected result type: %T", result6)
|
||||||
|
}
|
||||||
|
assert.Equal(t, "Updated Integration Collection", getResult2["name"])
|
||||||
|
t.Logf("✓ Step 6: Verified metadata update")
|
||||||
|
|
||||||
|
// Step 7: List collections (should include ours)
|
||||||
|
p7 := process.New("kb.collection.List").WithContext(ctx).WithAuthorized(authInfo)
|
||||||
|
p7.Args = []interface{}{
|
||||||
|
map[string]interface{}{
|
||||||
|
"page": 1,
|
||||||
|
"pagesize": 100,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
result7, err := p7.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
listResult := result7.(maps.MapStrAny)
|
||||||
|
assert.Contains(t, listResult, "data")
|
||||||
|
t.Logf("✓ Step 7: Listed collections")
|
||||||
|
|
||||||
|
// Step 8: Remove collection
|
||||||
|
p8 := process.New("kb.collection.Remove").WithContext(ctx).WithAuthorized(authInfo)
|
||||||
|
p8.Args = []interface{}{testCollectionID}
|
||||||
|
result8, err := p8.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
removeResult := result8.(maps.MapStrAny)
|
||||||
|
assert.Equal(t, true, removeResult["removed"])
|
||||||
|
t.Logf("✓ Step 8: Removed collection")
|
||||||
|
|
||||||
|
// Step 9: Verify collection no longer exists
|
||||||
|
p9 := process.New("kb.collection.Exists").WithContext(ctx)
|
||||||
|
p9.Args = []interface{}{testCollectionID}
|
||||||
|
result9, err := p9.Exec()
|
||||||
|
require.NoError(t, err)
|
||||||
|
existsResult3 := result9.(maps.MapStrAny)
|
||||||
|
assert.Equal(t, false, existsResult3["exists"])
|
||||||
|
t.Logf("✓ Step 9: Confirmed collection no longer exists")
|
||||||
|
|
||||||
|
t.Logf("✅ Full lifecycle completed successfully via Process API")
|
||||||
|
})
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue