diff --git a/kb/kb.go b/kb/kb.go index 2a7b6724..dd2de14d 100644 --- a/kb/kb.go +++ b/kb/kb.go @@ -23,7 +23,8 @@ var Instance types.GraphRag = nil // KnowledgeBase is the Knowledge Base instance type KnowledgeBase struct { - Config *kbtypes.Config // Knowledge Base configuration + Config *kbtypes.Config // Knowledge Base configuration + Providers *kbtypes.ProviderConfig // Multi-language provider configurations *graphrag.GraphRag } @@ -52,6 +53,13 @@ func Load(appConfig config.Config) (*KnowledgeBase, error) { return nil, err } + // Load providers from directories + providers, err := kbtypes.LoadProviders("kb") + if err != nil { + return nil, err + } + config.Providers = providers + // Set global configurations for providers to use kbtypes.SetGlobalPDF(config.PDF) kbtypes.SetGlobalFFmpeg(config.FFmpeg) @@ -69,7 +77,7 @@ func Load(appConfig config.Config) (*KnowledgeBase, error) { } // Set the instance - instance := &KnowledgeBase{Config: &config, GraphRag: graphRag} + instance := &KnowledgeBase{Config: &config, Providers: providers, GraphRag: graphRag} // Set the instance to the global variable Instance = instance @@ -88,48 +96,13 @@ func GetProviders(typ string, ids []string, locale string) ([]kbtypes.Provider, return nil, fmt.Errorf("knowledge base not initialized") } - // Get the configuration - conf := knowledgeBase.Config - if conf == nil { - return nil, fmt.Errorf("configuration not found") + // Default locale to "en" if empty + if locale == "" { + locale = "en" } - providers := []*kbtypes.Provider{} - switch typ { - case "chunking": - providers = conf.Chunkings - - case "converter": - providers = conf.Converters - - case "embedding": - providers = conf.Embeddings - - case "extractor": - providers = conf.Extractors - - case "fetcher": - providers = conf.Fetchers - - case "searcher": - providers = conf.Searchers - - case "reranker": - providers = conf.Rerankers - - case "vote": - providers = conf.Votes - - case "weight": - providers = conf.Weights - - case "score": - providers = conf.Scores - - default: - return nil, fmt.Errorf("invalid provider type: %s", typ) - - } + // Get providers for the requested type and language + providers := knowledgeBase.Providers.GetProviders(typ, locale) // Filter empty ids filteredIds := []string{} @@ -149,8 +122,13 @@ func GetProviders(typ string, ids []string, locale string) ([]kbtypes.Provider, return filteredProviders, nil } -// GetProvider returns a provider by id +// GetProvider returns a provider by id with default language "en" func GetProvider(typ string, id string) (*kbtypes.Provider, error) { + return GetProviderWithLanguage(typ, id, "en") +} + +// GetProviderWithLanguage returns a provider by id, type, and language +func GetProviderWithLanguage(typ string, id string, locale string) (*kbtypes.Provider, error) { if Instance == nil { return nil, fmt.Errorf("knowledge base not initialized") } @@ -160,53 +138,10 @@ func GetProvider(typ string, id string) (*kbtypes.Provider, error) { return nil, fmt.Errorf("knowledge base not initialized") } - conf := knowledgeBase.Config - if conf == nil { - return nil, fmt.Errorf("configuration not found") + // Default locale to "en" if empty + if locale == "" { + locale = "en" } - providers := []*kbtypes.Provider{} - switch typ { - case "chunking": - providers = conf.Chunkings - - case "converter": - providers = conf.Converters - - case "embedding": - providers = conf.Embeddings - - case "extractor": - providers = conf.Extractors - - case "fetcher": - providers = conf.Fetchers - - case "searcher": - providers = conf.Searchers - - case "reranker": - providers = conf.Rerankers - - case "vote": - providers = conf.Votes - - case "weight": - providers = conf.Weights - - case "score": - providers = conf.Scores - - default: - return nil, fmt.Errorf("invalid provider type: %s", typ) - } - - // Find the provider by id - for _, provider := range providers { - if provider.ID == id { - return provider, nil - } - } - - return nil, fmt.Errorf("provider %s not found", id) + return knowledgeBase.Providers.GetProvider(typ, id, locale) } diff --git a/kb/kb_test.go b/kb/kb_test.go index 6e534087..46c9e2b4 100644 --- a/kb/kb_test.go +++ b/kb/kb_test.go @@ -4,6 +4,7 @@ import ( "testing" "github.com/yaoapp/yao/config" + kbtypes "github.com/yaoapp/yao/kb/types" "github.com/yaoapp/yao/test" ) @@ -12,8 +13,134 @@ func TestLoad(t *testing.T) { test.Prepare(&testing.T{}, config.Conf) defer test.Clean() + kb, err := Load(config.Conf) + if err != nil { + t.Fatalf("Failed to load knowledge base: %v", err) + } + + // Test that providers are loaded + if kb != nil && kb.Providers != nil { + t.Logf("Knowledge base loaded successfully with providers") + } +} + +func TestGetProviders(t *testing.T) { + // Setup + test.Prepare(&testing.T{}, config.Conf) + defer test.Clean() + _, err := Load(config.Conf) if err != nil { t.Fatalf("Failed to load knowledge base: %v", err) } + + // Test getting providers for different languages + testCases := []struct { + providerType string + locale string + expectEmpty bool + }{ + {"chunking", "en", false}, + {"embedding", "en", false}, + {"chunking", "zh-cn", false}, + {"embedding", "zh-cn", false}, + {"chunking", "nonexistent", false}, // Should fallback to "en" + } + + for _, tc := range testCases { + providers, err := GetProviders(tc.providerType, []string{}, tc.locale) + if err != nil { + t.Errorf("Failed to get %s providers for locale %s: %v", tc.providerType, tc.locale, err) + continue + } + + if tc.expectEmpty && len(providers) > 0 { + t.Errorf("Expected empty providers for %s/%s, got %d", tc.providerType, tc.locale, len(providers)) + } else if !tc.expectEmpty && len(providers) == 0 { + t.Logf("No providers found for %s/%s (this may be expected if no provider files exist)", tc.providerType, tc.locale) + } else { + t.Logf("Found %d providers for %s/%s", len(providers), tc.providerType, tc.locale) + } + } +} + +func TestGetProviderWithLanguage(t *testing.T) { + // Setup + test.Prepare(&testing.T{}, config.Conf) + defer test.Clean() + + _, err := Load(config.Conf) + if err != nil { + t.Fatalf("Failed to load knowledge base: %v", err) + } + + // Test getting a specific provider with language + provider, err := GetProviderWithLanguage("chunking", "__yao.structured", "en") + if err != nil { + t.Logf("Provider __yao.structured not found for chunking/en: %v (this may be expected if provider files don't exist)", err) + } else { + t.Logf("Found provider: %s", provider.ID) + } + + // Test language fallback + provider, err = GetProviderWithLanguage("chunking", "__yao.structured", "nonexistent") + if err != nil { + t.Logf("Provider __yao.structured not found with fallback: %v (this may be expected if provider files don't exist)", err) + } else { + t.Logf("Found provider with fallback: %s", provider.ID) + } +} + +func TestLoadProviders(t *testing.T) { + // Setup + test.Prepare(&testing.T{}, config.Conf) + defer test.Clean() + + // Test loading providers from a directory + providers, err := kbtypes.LoadProviders("kb") + if err != nil { + t.Fatalf("Failed to load providers: %v", err) + } + + if providers == nil { + t.Fatal("Providers config is nil") + } + + // Check if provider maps are initialized + if providers.Chunkings == nil { + t.Error("Chunkings map is nil") + } + if providers.Embeddings == nil { + t.Error("Embeddings map is nil") + } + + t.Logf("Loaded providers successfully") +} + +func TestProviderConfigGetProviders(t *testing.T) { + // Setup + test.Prepare(&testing.T{}, config.Conf) + defer test.Clean() + + providers, err := kbtypes.LoadProviders("kb") + if err != nil { + t.Fatalf("Failed to load providers: %v", err) + } + + // Test getting providers for different types and languages + testCases := []string{"chunking", "embedding", "converter", "extractor", "fetcher"} + + for _, providerType := range testCases { + // Test with "en" + enProviders := providers.GetProviders(providerType, "en") + t.Logf("Found %d %s providers for 'en'", len(enProviders), providerType) + + // Test with "zh-cn" + zhProviders := providers.GetProviders(providerType, "zh-cn") + t.Logf("Found %d %s providers for 'zh-cn'", len(zhProviders), providerType) + + // Test with nonexistent language (should fallback to "en") + fallbackProviders := providers.GetProviders(providerType, "nonexistent") + t.Logf("Found %d %s providers for 'nonexistent' (fallback)", len(fallbackProviders), providerType) + } } diff --git a/kb/types/config.go b/kb/types/config.go index 51b86d7a..5f1406e5 100644 --- a/kb/types/config.go +++ b/kb/types/config.go @@ -272,20 +272,31 @@ func (c *Config) resolveAllEnvVars() error { // resolveProviderEnvVars resolves environment variables in provider configurations func (c *Config) resolveProviderEnvVars() error { - providerLists := [][]*Provider{ - c.Chunkings, c.Embeddings, c.Converters, c.Extractors, - c.Fetchers, c.Searchers, c.Rerankers, c.Votes, c.Weights, c.Scores, + if c.Providers == nil { + return nil } - for _, providers := range providerLists { - for _, provider := range providers { - for _, option := range provider.Options { - if option.Properties != nil { - resolved, err := c.resolveEnvVars(option.Properties) - if err != nil { - return err + // Resolve env vars for all provider types and languages + providerMaps := []map[string][]*Provider{ + c.Providers.Chunkings, c.Providers.Embeddings, c.Providers.Converters, c.Providers.Extractors, + c.Providers.Fetchers, c.Providers.Searchers, c.Providers.Rerankers, c.Providers.Votes, + c.Providers.Weights, c.Providers.Scores, + } + + for _, providerMap := range providerMaps { + if providerMap == nil { + continue + } + for _, providers := range providerMap { + for _, provider := range providers { + for _, option := range provider.Options { + if option.Properties != nil { + resolved, err := c.resolveEnvVars(option.Properties) + if err != nil { + return err + } + option.Properties = resolved } - option.Properties = resolved } } } @@ -335,8 +346,13 @@ func (c *Config) ComputeFeatures() Features { // File format support (based on converters) converterMap := make(map[string]bool) - for _, provider := range c.Converters { - converterMap[provider.ID] = true + if c.Providers != nil && c.Providers.Converters != nil { + // Check all languages for converter availability + for _, providers := range c.Providers.Converters { + for _, provider := range providers { + converterMap[provider.ID] = true + } + } } features.PlainText = true // Plain text is always supported as a basic feature @@ -346,13 +362,28 @@ func (c *Config) ComputeFeatures() Features { features.ImageAnalysis = converterMap["__yao.vision"] // Advanced features - features.EntityExtraction = len(c.Extractors) > 0 - features.WebFetching = len(c.Fetchers) > 0 - features.CustomSearch = len(c.Searchers) > 0 - features.ResultReranking = len(c.Rerankers) > 0 - features.SegmentVoting = len(c.Votes) > 0 - features.SegmentWeighting = len(c.Weights) > 0 - features.SegmentScoring = len(c.Scores) > 0 + if c.Providers != nil { + features.EntityExtraction = c.hasProvidersInAnyLanguage(c.Providers.Extractors) + features.WebFetching = c.hasProvidersInAnyLanguage(c.Providers.Fetchers) + features.CustomSearch = c.hasProvidersInAnyLanguage(c.Providers.Searchers) + features.ResultReranking = c.hasProvidersInAnyLanguage(c.Providers.Rerankers) + features.SegmentVoting = c.hasProvidersInAnyLanguage(c.Providers.Votes) + features.SegmentWeighting = c.hasProvidersInAnyLanguage(c.Providers.Weights) + features.SegmentScoring = c.hasProvidersInAnyLanguage(c.Providers.Scores) + } return features } + +// hasProvidersInAnyLanguage checks if there are providers available in any language +func (c *Config) hasProvidersInAnyLanguage(providerMap map[string][]*Provider) bool { + if providerMap == nil { + return false + } + for _, providers := range providerMap { + if len(providers) > 0 { + return true + } + } + return false +} diff --git a/kb/types/config_test.go b/kb/types/config_test.go index d2121e8c..35acf616 100644 --- a/kb/types/config_test.go +++ b/kb/types/config_test.go @@ -8,7 +8,7 @@ import ( "testing" ) -// Test data for configuration parsing +// Test data for configuration parsing (providers are now loaded from directories) const testConfigJSON = `{ "vector": { "driver": "qdrant", @@ -32,78 +32,14 @@ const testConfigJSON = `{ "ffmpeg_path": "/usr/bin/ffmpeg", "ffprobe_path": "/usr/bin/ffprobe", "enable_gpu": true - }, - "chunkings": [ - { - "id": "__yao.structured", - "label": "Document Structure", - "description": "Split text by document structure", - "default": true, - "options": [] - } - ], - "embeddings": [ - { - "id": "__yao.openai", - "label": "OpenAI", - "description": "OpenAI embeddings", - "default": true, - "options": [] - } - ], - "converters": [ - { - "id": "__yao.office", - "label": "Office Documents", - "description": "Process office documents", - "options": [] - }, - { - "id": "__yao.ocr", - "label": "OCR", - "description": "OCR processing", - "options": [] - } - ], - "extractors": [ - { - "id": "__yao.openai", - "label": "OpenAI Extractor", - "description": "Entity extraction", - "options": [] - } - ], - "fetchers": [ - { - "id": "__yao.http", - "label": "HTTP Fetcher", - "description": "Fetch from web", - "options": [] - } - ] + } }` const minimalConfigJSON = `{ "vector": { "driver": "qdrant", "config": {} - }, - "chunkings": [ - { - "id": "__yao.structured", - "label": "Document Structure", - "description": "Split text", - "options": [] - } - ], - "embeddings": [ - { - "id": "__yao.openai", - "label": "OpenAI", - "description": "OpenAI embeddings", - "options": [] - } - ] + } }` func TestParseConfigFromJSON(t *testing.T) { @@ -285,19 +221,37 @@ func TestConfig_ComputeFeatures(t *testing.T) { Graph: &GraphConfig{Driver: "neo4j"}, PDF: &PDFConfig{ConvertTool: "pdftoppm"}, FFmpeg: &FFmpegConfig{FFmpegPath: "/usr/bin/ffmpeg"}, - Converters: []*Provider{ - {ID: "__yao.office"}, - {ID: "__yao.ocr"}, - {ID: "__yao.whisper"}, - {ID: "__yao.vision"}, + Providers: &ProviderConfig{ + Converters: map[string][]*Provider{ + "en": { + {ID: "__yao.office"}, + {ID: "__yao.ocr"}, + {ID: "__yao.whisper"}, + {ID: "__yao.vision"}, + }, + }, + Extractors: map[string][]*Provider{ + "en": {{ID: "test"}}, + }, + Fetchers: map[string][]*Provider{ + "en": {{ID: "test"}}, + }, + Searchers: map[string][]*Provider{ + "en": {{ID: "test"}}, + }, + Rerankers: map[string][]*Provider{ + "en": {{ID: "test"}}, + }, + Votes: map[string][]*Provider{ + "en": {{ID: "test"}}, + }, + Weights: map[string][]*Provider{ + "en": {{ID: "test"}}, + }, + Scores: map[string][]*Provider{ + "en": {{ID: "test"}}, + }, }, - Extractors: []*Provider{{ID: "test"}}, - Fetchers: []*Provider{{ID: "test"}}, - Searchers: []*Provider{{ID: "test"}}, - Rerankers: []*Provider{{ID: "test"}}, - Votes: []*Provider{{ID: "test"}}, - Weights: []*Provider{{ID: "test"}}, - Scores: []*Provider{{ID: "test"}}, }, expected: Features{ GraphDatabase: true, @@ -320,9 +274,10 @@ func TestConfig_ComputeFeatures(t *testing.T) { { name: "minimal config", config: &Config{ - Graph: nil, - PDF: nil, - FFmpeg: nil, + Graph: nil, + PDF: nil, + FFmpeg: nil, + Providers: nil, }, expected: Features{ GraphDatabase: false, @@ -535,23 +490,7 @@ func TestConfig_ResolveEnvVarsOnParsing(t *testing.T) { "username": "$ENV.TEST_GRAPH_USER", "password": "$ENV.TEST_GRAPH_PASS" } - }, - "chunkings": [ - { - "id": "__yao.structured", - "label": "Document Structure", - "description": "Split text", - "options": [] - } - ], - "embeddings": [ - { - "id": "__yao.openai", - "label": "OpenAI", - "description": "OpenAI embeddings", - "options": [] - } - ] + } }` // Parse config from JSON @@ -582,3 +521,186 @@ func TestConfig_ResolveEnvVarsOnParsing(t *testing.T) { t.Errorf("Expected vector port to remain 6333.0, got %v (type %T)", config.Vector.Config["port"], config.Vector.Config["port"]) } } + +func TestProviderConfig_GetProviders(t *testing.T) { + // Create test provider config + providerConfig := &ProviderConfig{ + Chunkings: map[string][]*Provider{ + "en": { + {ID: "__yao.structured", Label: "Document Structure", Description: "Split by structure"}, + {ID: "__yao.semantic", Label: "Semantic Split", Description: "AI-powered splitting"}, + }, + "zh-cn": { + {ID: "__yao.structured", Label: "文档结构", Description: "按结构分割"}, + }, + }, + Embeddings: map[string][]*Provider{ + "en": { + {ID: "__yao.openai", Label: "OpenAI", Description: "OpenAI embeddings"}, + }, + }, + } + + tests := []struct { + name string + providerType string + language string + expectedLen int + expectedIDs []string + }{ + { + name: "get chunking providers for en", + providerType: "chunking", + language: "en", + expectedLen: 2, + expectedIDs: []string{"__yao.structured", "__yao.semantic"}, + }, + { + name: "get chunking providers for zh-cn", + providerType: "chunking", + language: "zh-cn", + expectedLen: 1, + expectedIDs: []string{"__yao.structured"}, + }, + { + name: "get embedding providers for en", + providerType: "embedding", + language: "en", + expectedLen: 1, + expectedIDs: []string{"__yao.openai"}, + }, + { + name: "fallback to en when language not found", + providerType: "embedding", + language: "fr", // Not available, should fallback to en + expectedLen: 1, + expectedIDs: []string{"__yao.openai"}, + }, + { + name: "return empty when provider type not found", + providerType: "nonexistent", + language: "en", + expectedLen: 0, + expectedIDs: []string{}, + }, + { + name: "return empty when no providers for language", + providerType: "converter", // Empty in test config + language: "en", + expectedLen: 0, + expectedIDs: []string{}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + providers := providerConfig.GetProviders(tt.providerType, tt.language) + + if len(providers) != tt.expectedLen { + t.Errorf("Expected %d providers, got %d", tt.expectedLen, len(providers)) + return + } + + // Check provider IDs + actualIDs := make([]string, len(providers)) + for i, provider := range providers { + actualIDs[i] = provider.ID + } + + for _, expectedID := range tt.expectedIDs { + found := false + for _, actualID := range actualIDs { + if actualID == expectedID { + found = true + break + } + } + if !found { + t.Errorf("Expected provider ID '%s' not found in results: %v", expectedID, actualIDs) + } + } + }) + } +} + +func TestProviderConfig_GetProvider(t *testing.T) { + // Create test provider config + providerConfig := &ProviderConfig{ + Chunkings: map[string][]*Provider{ + "en": { + {ID: "__yao.structured", Label: "Document Structure", Description: "Split by structure"}, + {ID: "__yao.semantic", Label: "Semantic Split", Description: "AI-powered splitting"}, + }, + "zh-cn": { + {ID: "__yao.structured", Label: "文档结构", Description: "按结构分割"}, + }, + }, + } + + tests := []struct { + name string + providerType string + providerID string + language string + expectError bool + expectedID string + }{ + { + name: "get existing provider in requested language", + providerType: "chunking", + providerID: "__yao.structured", + language: "en", + expectError: false, + expectedID: "__yao.structured", + }, + { + name: "get provider with language fallback", + providerType: "chunking", + providerID: "__yao.semantic", // Only exists in "en" + language: "fr", // Should fallback to "en" + expectError: false, + expectedID: "__yao.semantic", + }, + { + name: "provider not found", + providerType: "chunking", + providerID: "__yao.nonexistent", + language: "en", + expectError: true, + }, + { + name: "invalid provider type", + providerType: "invalid", + providerID: "__yao.structured", + language: "en", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + provider, err := providerConfig.GetProvider(tt.providerType, tt.providerID, tt.language) + + if tt.expectError { + if err == nil { + t.Error("Expected error, got nil") + } + return + } + + if err != nil { + t.Errorf("Unexpected error: %v", err) + return + } + + if provider == nil { + t.Error("Expected provider, got nil") + return + } + + if provider.ID != tt.expectedID { + t.Errorf("Expected provider ID '%s', got '%s'", tt.expectedID, provider.ID) + } + }) + } +} diff --git a/kb/types/provider.go b/kb/types/provider.go index 3257c7c0..3793d7f4 100644 --- a/kb/types/provider.go +++ b/kb/types/provider.go @@ -1,6 +1,14 @@ package types -import jsoniter "github.com/json-iterator/go" +import ( + "fmt" + "path/filepath" + "strings" + + jsoniter "github.com/json-iterator/go" + "github.com/yaoapp/gou/application" + "github.com/yaoapp/kun/log" +) // GetOption returns the option for a provider func (p *Provider) GetOption(id string) (*ProviderOption, bool) { @@ -41,3 +49,184 @@ func (p *ProviderOption) Parse(v interface{}) error { return nil } + +// LoadProviders loads providers from directories with language support +func LoadProviders(basePath string) (*ProviderConfig, error) { + config := &ProviderConfig{ + Chunkings: make(map[string][]*Provider), + Embeddings: make(map[string][]*Provider), + Converters: make(map[string][]*Provider), + Extractors: make(map[string][]*Provider), + Fetchers: make(map[string][]*Provider), + Searchers: make(map[string][]*Provider), + Rerankers: make(map[string][]*Provider), + Votes: make(map[string][]*Provider), + Weights: make(map[string][]*Provider), + Scores: make(map[string][]*Provider), + } + + // Provider type directories to load + providerTypes := []string{ + "chunkings", "embeddings", "converters", "extractions", + "fetchers", "searchers", "rerankers", "votes", "weights", "scores", + } + + for _, providerType := range providerTypes { + err := loadProviderType(basePath, providerType, config) + if err != nil { + log.Warn("[Knowledge Base] Failed to load %s providers: %v", providerType, err) + } + } + + return config, nil +} + +// loadProviderType loads providers for a specific type from language files +func loadProviderType(basePath, providerType string, config *ProviderConfig) error { + providerDir := filepath.Join(basePath, providerType) + + // Check if directory exists + exists, err := application.App.Exists(providerDir) + if err != nil { + return err + } + if !exists { + log.Debug("[Knowledge Base] Provider directory %s not found, skipping", providerDir) + return nil + } + + // Use Walk to find all provider files in the provider directory + err = application.App.Walk(providerDir, func(root, filename string, isdir bool) error { + if isdir { + return nil + } + + // Skip non-yao files + if !strings.HasSuffix(filename, ".yao") { + return nil + } + + // Extract language from filename (e.g., "en.yao" -> "en") + baseName := filepath.Base(filename) + language := strings.TrimSuffix(baseName, ".yao") + + // Load providers for this language + providers, err := loadProvidersForLanguage(providerDir, baseName) + if err != nil { + log.Warn("[Knowledge Base] Failed to load %s providers for language %s: %v", providerType, language, err) + return nil // Continue processing other files + } + + // Store providers in the appropriate map + switch providerType { + case "chunkings": + config.Chunkings[language] = providers + case "embeddings": + config.Embeddings[language] = providers + case "converters": + config.Converters[language] = providers + case "extractions": + config.Extractors[language] = providers + case "fetchers": + config.Fetchers[language] = providers + case "searchers": + config.Searchers[language] = providers + case "rerankers": + config.Rerankers[language] = providers + case "votes": + config.Votes[language] = providers + case "weights": + config.Weights[language] = providers + case "scores": + config.Scores[language] = providers + } + + log.Debug("[Knowledge Base] Loaded %d %s providers for language %s", len(providers), providerType, language) + return nil + }, "*.yao") + if err != nil { + return err + } + + return nil +} + +// loadProvidersForLanguage loads providers from a specific language file +func loadProvidersForLanguage(providerDir, filename string) ([]*Provider, error) { + filePath := filepath.Join(providerDir, filename) + + // Read the file + data, err := application.App.Read(filePath) + if err != nil { + return nil, err + } + + // Parse as array of providers + var providers []*Provider + err = application.Parse(filename, data, &providers) + if err != nil { + return nil, err + } + + return providers, nil +} + +// GetProviders returns providers for a specific type and language with fallback to "en" +func (pc *ProviderConfig) GetProviders(providerType, language string) []*Provider { + if pc == nil { + return []*Provider{} + } + + var providerMap map[string][]*Provider + switch providerType { + case "chunking": + providerMap = pc.Chunkings + case "embedding": + providerMap = pc.Embeddings + case "converter": + providerMap = pc.Converters + case "extractor": + providerMap = pc.Extractors + case "fetcher": + providerMap = pc.Fetchers + case "searcher": + providerMap = pc.Searchers + case "reranker": + providerMap = pc.Rerankers + case "vote": + providerMap = pc.Votes + case "weight": + providerMap = pc.Weights + case "score": + providerMap = pc.Scores + default: + return []*Provider{} + } + + // Try to get providers for the requested language + if providers, exists := providerMap[language]; exists && len(providers) > 0 { + return providers + } + + // Fallback to "en" if requested language not found + if language != "en" { + if providers, exists := providerMap["en"]; exists && len(providers) > 0 { + return providers + } + } + + return []*Provider{} +} + +// GetProvider returns a specific provider by ID, type, and language with fallback to "en" +func (pc *ProviderConfig) GetProvider(providerType, id, language string) (*Provider, error) { + providers := pc.GetProviders(providerType, language) + + for _, provider := range providers { + if provider.ID == id { + return provider, nil + } + } + + return nil, fmt.Errorf("provider %s not found for type %s and language %s", id, providerType, language) +} diff --git a/kb/types/types.go b/kb/types/types.go index ecba3793..4011d2b4 100644 --- a/kb/types/types.go +++ b/kb/types/types.go @@ -75,22 +75,28 @@ type Config struct { // Concurrency limits for task processing (Optional) Limits *LimitsConfig `json:"limits,omitempty" yaml:"limits,omitempty"` - // Provider configurations - Chunkings []*Provider `json:"chunkings" yaml:"chunkings"` // Text splitting providers (Required - at least one) - Embeddings []*Provider `json:"embeddings" yaml:"embeddings"` // Text vectorization providers (Required - at least one) - Converters []*Provider `json:"converters,omitempty" yaml:"converters,omitempty"` // File processing converters (Optional) - Extractors []*Provider `json:"extractors,omitempty" yaml:"extractors,omitempty"` // Entity and relationship extractors (Optional) - Fetchers []*Provider `json:"fetchers,omitempty" yaml:"fetchers,omitempty"` // File fetchers (Optional) - Searchers []*Provider `json:"searchers,omitempty" yaml:"searchers,omitempty"` // Search providers (Optional) - Rerankers []*Provider `json:"rerankers,omitempty" yaml:"rerankers,omitempty"` // Reranking providers (Optional) - Votes []*Provider `json:"votes,omitempty" yaml:"votes,omitempty"` // Voting providers (Optional) - Weights []*Provider `json:"weights,omitempty" yaml:"weights,omitempty"` // Weighting providers (Optional) - Scores []*Provider `json:"scores,omitempty" yaml:"scores,omitempty"` // Scoring providers (Optional) + // Multi-language provider configurations (loaded from directories) + Providers *ProviderConfig `json:"-"` // Loaded from provider directories, not serialized // Feature flags (computed during parsing, not serialized) Features Features `json:"-"` } +// ProviderConfig holds providers organized by language +type ProviderConfig struct { + // Provider configurations by language (e.g., "en", "zh-cn") + Chunkings map[string][]*Provider `json:"-"` // Text splitting providers by language + Embeddings map[string][]*Provider `json:"-"` // Text vectorization providers by language + Converters map[string][]*Provider `json:"-"` // File processing converters by language + Extractors map[string][]*Provider `json:"-"` // Entity and relationship extractors by language + Fetchers map[string][]*Provider `json:"-"` // File fetchers by language + Searchers map[string][]*Provider `json:"-"` // Search providers by language + Rerankers map[string][]*Provider `json:"-"` // Reranking providers by language + Votes map[string][]*Provider `json:"-"` // Voting providers by language + Weights map[string][]*Provider `json:"-"` // Weighting providers by language + Scores map[string][]*Provider `json:"-"` // Scoring providers by language +} + // VectorConfig represents vector database configuration type VectorConfig struct { Driver string `json:"driver" yaml:"driver"` // Required, currently only support "qdrant" diff --git a/openapi/kb/provider.go b/openapi/kb/provider.go index 069ee137..1f741b99 100644 --- a/openapi/kb/provider.go +++ b/openapi/kb/provider.go @@ -12,7 +12,7 @@ import ( // GetProviders get all providers func GetProviders(c *gin.Context) { providerType := c.Param("providerType") - locale := c.Query("locale") + locale := strings.ToLower(c.Query("locale")) if locale == "" { locale = "en" } @@ -48,7 +48,7 @@ func GetProviders(c *gin.Context) { func GetProviderSchema(c *gin.Context) { providerType := c.Param("providerType") providerID := c.Param("providerID") - locale := c.Query("locale") + locale := strings.ToLower(c.Query("locale")) if locale == "" { locale = "en" } @@ -62,7 +62,7 @@ func GetProviderSchema(c *gin.Context) { return } - provider, err := kb.GetProvider(providerType, providerID) + provider, err := kb.GetProviderWithLanguage(providerType, providerID, locale) if err != nil { errorResp := &response.ErrorResponse{ Code: response.ErrServerError.Code, diff --git a/openapi/kb/utils.go b/openapi/kb/utils.go index 3321b421..0bd10531 100644 --- a/openapi/kb/utils.go +++ b/openapi/kb/utils.go @@ -15,13 +15,14 @@ Usage Examples: 1. AddFile API (converter will be auto-detected based on file info): { "collection_id": "my_collection", + "locale": "en", "file_id": "uploaded_file_123", "chunking": { - "provider_id": "text_splitter", - "option_id": "default" + "provider_id": "__yao.structured", + "option_id": "standard" }, "embedding": { - "provider_id": "openai", + "provider_id": "__yao.openai", "option_id": "text-embedding-3-small" }, "doc_id": "document_001", @@ -30,34 +31,39 @@ Usage Examples: } } -2. AddText API: +2. AddText API with Chinese locale: { "collection_id": "my_collection", - "text": "This is the text content to be processed.", + "locale": "zh-cn", + "text": "这是要处理的文本内容。", "chunking": { - "provider_id": "text_splitter" + "provider_id": "__yao.structured" }, "embedding": { - "provider_id": "openai" + "provider_id": "__yao.fastembed", + "option_id": "fastembed-chinese" } } 3. AddSegments API: { "collection_id": "my_collection", + "locale": "en", "doc_id": "document_001", "segment_texts": [ {"text": "First segment", "metadata": {"page": 1}}, {"text": "Second segment", "metadata": {"page": 2}} ], "embedding": { - "provider_id": "openai", + "provider_id": "__yao.openai", "option_id": "text-embedding-3-small" } } Note: +- If no locale is specified, defaults to "en" - If no option_id is specified, the default option from provider configuration will be selected +- Providers are loaded based on locale with fallback to "en" if the specified locale is not available - For AddFile API, converter will be auto-detected based on filename and content_type obtained from GetFileInfo(file_id) - ToUpsertOptions() can be called without parameters, or with filename and contentType for converter auto-detection */ @@ -76,6 +82,9 @@ type BaseUpsertRequest struct { // Collection ID - this will be mapped to UpsertOptions.CollectionID CollectionID string `json:"collection_id" binding:"required"` + // Language/locale for provider selection (defaults to "en") + Locale string `json:"locale,omitempty"` + // Provider configurations Chunking *ProviderConfig `json:"chunking" binding:"required"` Embedding *ProviderConfig `json:"embedding" binding:"required"` @@ -123,7 +132,7 @@ type UpdateSegmentsRequest struct { // If OptionID is provided, it looks up the option from the provider // If Option is provided directly, it uses the Option field // If neither is provided, it selects the default option from provider's Options -func resolveProviderOption(config *ProviderConfig) (*kbtypes.ProviderOption, error) { +func resolveProviderOption(config *ProviderConfig, locale string) (*kbtypes.ProviderOption, error) { if config == nil { return nil, fmt.Errorf("provider config is required") } @@ -142,20 +151,20 @@ func resolveProviderOption(config *ProviderConfig) (*kbtypes.ProviderOption, err return nil, fmt.Errorf("KB instance is not initialized") } - // Find the provider in KB config - var provider *kbtypes.Provider - kbConfig := kb.Instance.(*kb.KnowledgeBase).Config - - // Check all provider types to find the matching provider - allProviders := [][]*kbtypes.Provider{ - kbConfig.Chunkings, - kbConfig.Embeddings, - kbConfig.Converters, - kbConfig.Extractors, - kbConfig.Fetchers, + // Default locale to "en" if not provided + if locale == "" { + locale = "en" } - for _, providers := range allProviders { + // Find the provider using the new multi-language system + var provider *kbtypes.Provider + kbInstance := kb.Instance.(*kb.KnowledgeBase) + + // Check all provider types to find the matching provider + providerTypes := []string{"chunking", "embedding", "converter", "extractor", "fetcher"} + + for _, providerType := range providerTypes { + providers := kbInstance.Providers.GetProviders(providerType, locale) for _, p := range providers { if p.ID == config.ProviderID { provider = p @@ -168,7 +177,7 @@ func resolveProviderOption(config *ProviderConfig) (*kbtypes.ProviderOption, err } if provider == nil { - return nil, fmt.Errorf("provider %s not found", config.ProviderID) + return nil, fmt.Errorf("provider %s not found for locale %s", config.ProviderID, locale) } // If OptionID is provided, look it up from the provider @@ -207,6 +216,12 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp contentType = fileInfo[1] } + // Default locale to "en" if not specified + locale := r.Locale + if locale == "" { + locale = "en" + } + options := &types.UpsertOptions{ CollectionID: r.CollectionID, // Collection ID maps to CollectionID DocID: r.DocID, @@ -214,7 +229,7 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp } // Resolve and create chunking provider - chunkingOption, err := resolveProviderOption(r.Chunking) + chunkingOption, err := resolveProviderOption(r.Chunking, locale) if err != nil { return nil, fmt.Errorf("failed to resolve chunking provider: %w", err) } @@ -233,7 +248,7 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp options.ChunkingOptions = chunkingOpts // Resolve and create embedding provider - embeddingOption, err := resolveProviderOption(r.Embedding) + embeddingOption, err := resolveProviderOption(r.Embedding, locale) if err != nil { return nil, fmt.Errorf("failed to resolve embedding provider: %w", err) } @@ -246,7 +261,7 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp // Optional providers if r.Extraction != nil { - extractionOption, err := resolveProviderOption(r.Extraction) + extractionOption, err := resolveProviderOption(r.Extraction, locale) if err != nil { return nil, fmt.Errorf("failed to resolve extraction provider: %w", err) } @@ -259,7 +274,7 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp } if r.Fetcher != nil { - fetcherOption, err := resolveProviderOption(r.Fetcher) + fetcherOption, err := resolveProviderOption(r.Fetcher, locale) if err != nil { return nil, fmt.Errorf("failed to resolve fetcher provider: %w", err) } @@ -274,7 +289,7 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp // Handle converter - auto-detect if not specified if r.Converter != nil { // User specified converter - converterOption, err := resolveProviderOption(r.Converter) + converterOption, err := resolveProviderOption(r.Converter, locale) if err != nil { return nil, fmt.Errorf("failed to resolve converter provider: %w", err) } @@ -297,7 +312,7 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp ProviderID: converterID, } - converterOption, err := resolveProviderOption(converterConfig) + converterOption, err := resolveProviderOption(converterConfig, locale) if err != nil { return nil, fmt.Errorf("failed to resolve auto-detected converter provider: %w", err) } diff --git a/widgets/app/app.go b/widgets/app/app.go index 47380003..7f82d4cc 100644 --- a/widgets/app/app.go +++ b/widgets/app/app.go @@ -22,6 +22,7 @@ import ( "github.com/yaoapp/yao/data" "github.com/yaoapp/yao/i18n" "github.com/yaoapp/yao/kb" + kbtypes "github.com/yaoapp/yao/kb/types" "github.com/yaoapp/yao/neo" "github.com/yaoapp/yao/neo/assistant" "github.com/yaoapp/yao/openapi" @@ -599,74 +600,62 @@ func processXgen(process *process.Process) interface{} { kbConfig := map[string]interface{}{} if kb.Instance != nil { if knowledgebase, ok := kb.Instance.(*kb.KnowledgeBase); ok && knowledgebase.Config != nil { - chunkings := []string{} - if knowledgebase.Config.Chunkings != nil { - for _, chunking := range knowledgebase.Config.Chunkings { - chunkings = append(chunkings, chunking.ID) - } + // Use the current language setting for provider selection + currentLang := lang + if currentLang == "" { + currentLang = "en" // Default to English } - embeddings := []string{} - if knowledgebase.Config.Embeddings != nil { - for _, embedding := range knowledgebase.Config.Embeddings { - embeddings = append(embeddings, embedding.ID) + // Helper function to extract provider IDs from multi-language providers + extractProviderIDs := func(providerMap map[string][]*kbtypes.Provider) []string { + ids := []string{} + if providerMap == nil { + return ids } + + // Try current language first + if providers, exists := providerMap[currentLang]; exists { + for _, provider := range providers { + ids = append(ids, provider.ID) + } + return ids + } + + // Fallback to English + if currentLang != "en" { + if providers, exists := providerMap["en"]; exists { + for _, provider := range providers { + ids = append(ids, provider.ID) + } + return ids + } + } + + // If no providers found for current language or English, return all available + for _, providers := range providerMap { + for _, provider := range providers { + ids = append(ids, provider.ID) + } + break // Just take the first available language + } + + return ids } - converters := []string{} - if knowledgebase.Config.Converters != nil { - for _, converter := range knowledgebase.Config.Converters { - converters = append(converters, converter.ID) - } - } + var chunkings, embeddings, converters, extractors, fetchers []string + var searchers, rerankers, votes, weights, scores []string - extractors := []string{} - if knowledgebase.Config.Extractors != nil { - for _, extractor := range knowledgebase.Config.Extractors { - extractors = append(extractors, extractor.ID) - } - } - - fetchers := []string{} - if knowledgebase.Config.Fetchers != nil { - for _, fetcher := range knowledgebase.Config.Fetchers { - fetchers = append(fetchers, fetcher.ID) - } - } - - searchers := []string{} - if knowledgebase.Config.Searchers != nil { - for _, searcher := range knowledgebase.Config.Searchers { - searchers = append(searchers, searcher.ID) - } - } - - rerankers := []string{} - if knowledgebase.Config.Rerankers != nil { - for _, reranker := range knowledgebase.Config.Rerankers { - rerankers = append(rerankers, reranker.ID) - } - } - - votes := []string{} - if knowledgebase.Config.Votes != nil { - for _, vote := range knowledgebase.Config.Votes { - votes = append(votes, vote.ID) - } - } - - weights := []string{} - if knowledgebase.Config.Weights != nil { - for _, weight := range knowledgebase.Config.Weights { - weights = append(weights, weight.ID) - } - } - - scores := []string{} - if knowledgebase.Config.Scores != nil { - for _, score := range knowledgebase.Config.Scores { - scores = append(scores, score.ID) - } + if knowledgebase.Providers != nil { + chunkings = extractProviderIDs(knowledgebase.Providers.Chunkings) + embeddings = extractProviderIDs(knowledgebase.Providers.Embeddings) + converters = extractProviderIDs(knowledgebase.Providers.Converters) + extractors = extractProviderIDs(knowledgebase.Providers.Extractors) + fetchers = extractProviderIDs(knowledgebase.Providers.Fetchers) + searchers = extractProviderIDs(knowledgebase.Providers.Searchers) + rerankers = extractProviderIDs(knowledgebase.Providers.Rerankers) + votes = extractProviderIDs(knowledgebase.Providers.Votes) + weights = extractProviderIDs(knowledgebase.Providers.Weights) + scores = extractProviderIDs(knowledgebase.Providers.Scores) } kbConfig = map[string]interface{}{