Merge pull request #1093 from trheyi/main

Enhance knowledge base provider management with multi-language support
This commit is contained in:
Max 2025-08-09 16:58:33 +08:00 committed by GitHub
commit 7af1eae1fb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 728 additions and 314 deletions

115
kb/kb.go
View file

@ -23,7 +23,8 @@ var Instance types.GraphRag = nil
// KnowledgeBase is the Knowledge Base instance // KnowledgeBase is the Knowledge Base instance
type KnowledgeBase struct { type KnowledgeBase struct {
Config *kbtypes.Config // Knowledge Base configuration Config *kbtypes.Config // Knowledge Base configuration
Providers *kbtypes.ProviderConfig // Multi-language provider configurations
*graphrag.GraphRag *graphrag.GraphRag
} }
@ -52,6 +53,13 @@ func Load(appConfig config.Config) (*KnowledgeBase, error) {
return nil, err 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 // Set global configurations for providers to use
kbtypes.SetGlobalPDF(config.PDF) kbtypes.SetGlobalPDF(config.PDF)
kbtypes.SetGlobalFFmpeg(config.FFmpeg) kbtypes.SetGlobalFFmpeg(config.FFmpeg)
@ -69,7 +77,7 @@ func Load(appConfig config.Config) (*KnowledgeBase, error) {
} }
// Set the instance // Set the instance
instance := &KnowledgeBase{Config: &config, GraphRag: graphRag} instance := &KnowledgeBase{Config: &config, Providers: providers, GraphRag: graphRag}
// Set the instance to the global variable // Set the instance to the global variable
Instance = instance Instance = instance
@ -88,48 +96,13 @@ func GetProviders(typ string, ids []string, locale string) ([]kbtypes.Provider,
return nil, fmt.Errorf("knowledge base not initialized") return nil, fmt.Errorf("knowledge base not initialized")
} }
// Get the configuration // Default locale to "en" if empty
conf := knowledgeBase.Config if locale == "" {
if conf == nil { locale = "en"
return nil, fmt.Errorf("configuration not found")
} }
providers := []*kbtypes.Provider{} // Get providers for the requested type and language
switch typ { providers := knowledgeBase.Providers.GetProviders(typ, locale)
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)
}
// Filter empty ids // Filter empty ids
filteredIds := []string{} filteredIds := []string{}
@ -149,8 +122,13 @@ func GetProviders(typ string, ids []string, locale string) ([]kbtypes.Provider,
return filteredProviders, nil 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) { 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 { if Instance == nil {
return nil, fmt.Errorf("knowledge base not initialized") 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") return nil, fmt.Errorf("knowledge base not initialized")
} }
conf := knowledgeBase.Config // Default locale to "en" if empty
if conf == nil { if locale == "" {
return nil, fmt.Errorf("configuration not found") locale = "en"
} }
providers := []*kbtypes.Provider{} return knowledgeBase.Providers.GetProvider(typ, id, locale)
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)
} }

View file

@ -4,6 +4,7 @@ import (
"testing" "testing"
"github.com/yaoapp/yao/config" "github.com/yaoapp/yao/config"
kbtypes "github.com/yaoapp/yao/kb/types"
"github.com/yaoapp/yao/test" "github.com/yaoapp/yao/test"
) )
@ -12,8 +13,134 @@ func TestLoad(t *testing.T) {
test.Prepare(&testing.T{}, config.Conf) test.Prepare(&testing.T{}, config.Conf)
defer test.Clean() 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) _, err := Load(config.Conf)
if err != nil { if err != nil {
t.Fatalf("Failed to load knowledge base: %v", err) 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)
}
} }

View file

@ -272,20 +272,31 @@ func (c *Config) resolveAllEnvVars() error {
// resolveProviderEnvVars resolves environment variables in provider configurations // resolveProviderEnvVars resolves environment variables in provider configurations
func (c *Config) resolveProviderEnvVars() error { func (c *Config) resolveProviderEnvVars() error {
providerLists := [][]*Provider{ if c.Providers == nil {
c.Chunkings, c.Embeddings, c.Converters, c.Extractors, return nil
c.Fetchers, c.Searchers, c.Rerankers, c.Votes, c.Weights, c.Scores,
} }
for _, providers := range providerLists { // Resolve env vars for all provider types and languages
for _, provider := range providers { providerMaps := []map[string][]*Provider{
for _, option := range provider.Options { c.Providers.Chunkings, c.Providers.Embeddings, c.Providers.Converters, c.Providers.Extractors,
if option.Properties != nil { c.Providers.Fetchers, c.Providers.Searchers, c.Providers.Rerankers, c.Providers.Votes,
resolved, err := c.resolveEnvVars(option.Properties) c.Providers.Weights, c.Providers.Scores,
if err != nil { }
return err
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) // File format support (based on converters)
converterMap := make(map[string]bool) converterMap := make(map[string]bool)
for _, provider := range c.Converters { if c.Providers != nil && c.Providers.Converters != nil {
converterMap[provider.ID] = true // 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 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"] features.ImageAnalysis = converterMap["__yao.vision"]
// Advanced features // Advanced features
features.EntityExtraction = len(c.Extractors) > 0 if c.Providers != nil {
features.WebFetching = len(c.Fetchers) > 0 features.EntityExtraction = c.hasProvidersInAnyLanguage(c.Providers.Extractors)
features.CustomSearch = len(c.Searchers) > 0 features.WebFetching = c.hasProvidersInAnyLanguage(c.Providers.Fetchers)
features.ResultReranking = len(c.Rerankers) > 0 features.CustomSearch = c.hasProvidersInAnyLanguage(c.Providers.Searchers)
features.SegmentVoting = len(c.Votes) > 0 features.ResultReranking = c.hasProvidersInAnyLanguage(c.Providers.Rerankers)
features.SegmentWeighting = len(c.Weights) > 0 features.SegmentVoting = c.hasProvidersInAnyLanguage(c.Providers.Votes)
features.SegmentScoring = len(c.Scores) > 0 features.SegmentWeighting = c.hasProvidersInAnyLanguage(c.Providers.Weights)
features.SegmentScoring = c.hasProvidersInAnyLanguage(c.Providers.Scores)
}
return features 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
}

View file

@ -8,7 +8,7 @@ import (
"testing" "testing"
) )
// Test data for configuration parsing // Test data for configuration parsing (providers are now loaded from directories)
const testConfigJSON = `{ const testConfigJSON = `{
"vector": { "vector": {
"driver": "qdrant", "driver": "qdrant",
@ -32,78 +32,14 @@ const testConfigJSON = `{
"ffmpeg_path": "/usr/bin/ffmpeg", "ffmpeg_path": "/usr/bin/ffmpeg",
"ffprobe_path": "/usr/bin/ffprobe", "ffprobe_path": "/usr/bin/ffprobe",
"enable_gpu": true "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 = `{ const minimalConfigJSON = `{
"vector": { "vector": {
"driver": "qdrant", "driver": "qdrant",
"config": {} "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) { func TestParseConfigFromJSON(t *testing.T) {
@ -285,19 +221,37 @@ func TestConfig_ComputeFeatures(t *testing.T) {
Graph: &GraphConfig{Driver: "neo4j"}, Graph: &GraphConfig{Driver: "neo4j"},
PDF: &PDFConfig{ConvertTool: "pdftoppm"}, PDF: &PDFConfig{ConvertTool: "pdftoppm"},
FFmpeg: &FFmpegConfig{FFmpegPath: "/usr/bin/ffmpeg"}, FFmpeg: &FFmpegConfig{FFmpegPath: "/usr/bin/ffmpeg"},
Converters: []*Provider{ Providers: &ProviderConfig{
{ID: "__yao.office"}, Converters: map[string][]*Provider{
{ID: "__yao.ocr"}, "en": {
{ID: "__yao.whisper"}, {ID: "__yao.office"},
{ID: "__yao.vision"}, {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{ expected: Features{
GraphDatabase: true, GraphDatabase: true,
@ -320,9 +274,10 @@ func TestConfig_ComputeFeatures(t *testing.T) {
{ {
name: "minimal config", name: "minimal config",
config: &Config{ config: &Config{
Graph: nil, Graph: nil,
PDF: nil, PDF: nil,
FFmpeg: nil, FFmpeg: nil,
Providers: nil,
}, },
expected: Features{ expected: Features{
GraphDatabase: false, GraphDatabase: false,
@ -535,23 +490,7 @@ func TestConfig_ResolveEnvVarsOnParsing(t *testing.T) {
"username": "$ENV.TEST_GRAPH_USER", "username": "$ENV.TEST_GRAPH_USER",
"password": "$ENV.TEST_GRAPH_PASS" "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 // 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"]) 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)
}
})
}
}

View file

@ -1,6 +1,14 @@
package types 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 // GetOption returns the option for a provider
func (p *Provider) GetOption(id string) (*ProviderOption, bool) { func (p *Provider) GetOption(id string) (*ProviderOption, bool) {
@ -41,3 +49,184 @@ func (p *ProviderOption) Parse(v interface{}) error {
return nil 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)
}

View file

@ -75,22 +75,28 @@ type Config struct {
// Concurrency limits for task processing (Optional) // Concurrency limits for task processing (Optional)
Limits *LimitsConfig `json:"limits,omitempty" yaml:"limits,omitempty"` Limits *LimitsConfig `json:"limits,omitempty" yaml:"limits,omitempty"`
// Provider configurations // Multi-language provider configurations (loaded from directories)
Chunkings []*Provider `json:"chunkings" yaml:"chunkings"` // Text splitting providers (Required - at least one) Providers *ProviderConfig `json:"-"` // Loaded from provider directories, not serialized
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)
// Feature flags (computed during parsing, not serialized) // Feature flags (computed during parsing, not serialized)
Features Features `json:"-"` 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 // VectorConfig represents vector database configuration
type VectorConfig struct { type VectorConfig struct {
Driver string `json:"driver" yaml:"driver"` // Required, currently only support "qdrant" Driver string `json:"driver" yaml:"driver"` // Required, currently only support "qdrant"

View file

@ -12,7 +12,7 @@ import (
// GetProviders get all providers // GetProviders get all providers
func GetProviders(c *gin.Context) { func GetProviders(c *gin.Context) {
providerType := c.Param("providerType") providerType := c.Param("providerType")
locale := c.Query("locale") locale := strings.ToLower(c.Query("locale"))
if locale == "" { if locale == "" {
locale = "en" locale = "en"
} }
@ -48,7 +48,7 @@ func GetProviders(c *gin.Context) {
func GetProviderSchema(c *gin.Context) { func GetProviderSchema(c *gin.Context) {
providerType := c.Param("providerType") providerType := c.Param("providerType")
providerID := c.Param("providerID") providerID := c.Param("providerID")
locale := c.Query("locale") locale := strings.ToLower(c.Query("locale"))
if locale == "" { if locale == "" {
locale = "en" locale = "en"
} }
@ -62,7 +62,7 @@ func GetProviderSchema(c *gin.Context) {
return return
} }
provider, err := kb.GetProvider(providerType, providerID) provider, err := kb.GetProviderWithLanguage(providerType, providerID, locale)
if err != nil { if err != nil {
errorResp := &response.ErrorResponse{ errorResp := &response.ErrorResponse{
Code: response.ErrServerError.Code, Code: response.ErrServerError.Code,

View file

@ -15,13 +15,14 @@ Usage Examples:
1. AddFile API (converter will be auto-detected based on file info): 1. AddFile API (converter will be auto-detected based on file info):
{ {
"collection_id": "my_collection", "collection_id": "my_collection",
"locale": "en",
"file_id": "uploaded_file_123", "file_id": "uploaded_file_123",
"chunking": { "chunking": {
"provider_id": "text_splitter", "provider_id": "__yao.structured",
"option_id": "default" "option_id": "standard"
}, },
"embedding": { "embedding": {
"provider_id": "openai", "provider_id": "__yao.openai",
"option_id": "text-embedding-3-small" "option_id": "text-embedding-3-small"
}, },
"doc_id": "document_001", "doc_id": "document_001",
@ -30,34 +31,39 @@ Usage Examples:
} }
} }
2. AddText API: 2. AddText API with Chinese locale:
{ {
"collection_id": "my_collection", "collection_id": "my_collection",
"text": "This is the text content to be processed.", "locale": "zh-cn",
"text": "这是要处理的文本内容。",
"chunking": { "chunking": {
"provider_id": "text_splitter" "provider_id": "__yao.structured"
}, },
"embedding": { "embedding": {
"provider_id": "openai" "provider_id": "__yao.fastembed",
"option_id": "fastembed-chinese"
} }
} }
3. AddSegments API: 3. AddSegments API:
{ {
"collection_id": "my_collection", "collection_id": "my_collection",
"locale": "en",
"doc_id": "document_001", "doc_id": "document_001",
"segment_texts": [ "segment_texts": [
{"text": "First segment", "metadata": {"page": 1}}, {"text": "First segment", "metadata": {"page": 1}},
{"text": "Second segment", "metadata": {"page": 2}} {"text": "Second segment", "metadata": {"page": 2}}
], ],
"embedding": { "embedding": {
"provider_id": "openai", "provider_id": "__yao.openai",
"option_id": "text-embedding-3-small" "option_id": "text-embedding-3-small"
} }
} }
Note: Note:
- If no locale is specified, defaults to "en"
- If no option_id is specified, the default option from provider configuration will be selected - 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) - 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 - 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 // Collection ID - this will be mapped to UpsertOptions.CollectionID
CollectionID string `json:"collection_id" binding:"required"` CollectionID string `json:"collection_id" binding:"required"`
// Language/locale for provider selection (defaults to "en")
Locale string `json:"locale,omitempty"`
// Provider configurations // Provider configurations
Chunking *ProviderConfig `json:"chunking" binding:"required"` Chunking *ProviderConfig `json:"chunking" binding:"required"`
Embedding *ProviderConfig `json:"embedding" 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 OptionID is provided, it looks up the option from the provider
// If Option is provided directly, it uses the Option field // If Option is provided directly, it uses the Option field
// If neither is provided, it selects the default option from provider's Options // 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 { if config == nil {
return nil, fmt.Errorf("provider config is required") 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") return nil, fmt.Errorf("KB instance is not initialized")
} }
// Find the provider in KB config // Default locale to "en" if not provided
var provider *kbtypes.Provider if locale == "" {
kbConfig := kb.Instance.(*kb.KnowledgeBase).Config locale = "en"
// Check all provider types to find the matching provider
allProviders := [][]*kbtypes.Provider{
kbConfig.Chunkings,
kbConfig.Embeddings,
kbConfig.Converters,
kbConfig.Extractors,
kbConfig.Fetchers,
} }
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 { for _, p := range providers {
if p.ID == config.ProviderID { if p.ID == config.ProviderID {
provider = p provider = p
@ -168,7 +177,7 @@ func resolveProviderOption(config *ProviderConfig) (*kbtypes.ProviderOption, err
} }
if provider == nil { 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 // 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] contentType = fileInfo[1]
} }
// Default locale to "en" if not specified
locale := r.Locale
if locale == "" {
locale = "en"
}
options := &types.UpsertOptions{ options := &types.UpsertOptions{
CollectionID: r.CollectionID, // Collection ID maps to CollectionID CollectionID: r.CollectionID, // Collection ID maps to CollectionID
DocID: r.DocID, DocID: r.DocID,
@ -214,7 +229,7 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp
} }
// Resolve and create chunking provider // Resolve and create chunking provider
chunkingOption, err := resolveProviderOption(r.Chunking) chunkingOption, err := resolveProviderOption(r.Chunking, locale)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to resolve chunking provider: %w", err) 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 options.ChunkingOptions = chunkingOpts
// Resolve and create embedding provider // Resolve and create embedding provider
embeddingOption, err := resolveProviderOption(r.Embedding) embeddingOption, err := resolveProviderOption(r.Embedding, locale)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to resolve embedding provider: %w", err) 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 // Optional providers
if r.Extraction != nil { if r.Extraction != nil {
extractionOption, err := resolveProviderOption(r.Extraction) extractionOption, err := resolveProviderOption(r.Extraction, locale)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to resolve extraction provider: %w", err) 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 { if r.Fetcher != nil {
fetcherOption, err := resolveProviderOption(r.Fetcher) fetcherOption, err := resolveProviderOption(r.Fetcher, locale)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to resolve fetcher provider: %w", err) 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 // Handle converter - auto-detect if not specified
if r.Converter != nil { if r.Converter != nil {
// User specified converter // User specified converter
converterOption, err := resolveProviderOption(r.Converter) converterOption, err := resolveProviderOption(r.Converter, locale)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to resolve converter provider: %w", err) 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, ProviderID: converterID,
} }
converterOption, err := resolveProviderOption(converterConfig) converterOption, err := resolveProviderOption(converterConfig, locale)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to resolve auto-detected converter provider: %w", err) return nil, fmt.Errorf("failed to resolve auto-detected converter provider: %w", err)
} }

View file

@ -22,6 +22,7 @@ import (
"github.com/yaoapp/yao/data" "github.com/yaoapp/yao/data"
"github.com/yaoapp/yao/i18n" "github.com/yaoapp/yao/i18n"
"github.com/yaoapp/yao/kb" "github.com/yaoapp/yao/kb"
kbtypes "github.com/yaoapp/yao/kb/types"
"github.com/yaoapp/yao/neo" "github.com/yaoapp/yao/neo"
"github.com/yaoapp/yao/neo/assistant" "github.com/yaoapp/yao/neo/assistant"
"github.com/yaoapp/yao/openapi" "github.com/yaoapp/yao/openapi"
@ -599,74 +600,62 @@ func processXgen(process *process.Process) interface{} {
kbConfig := map[string]interface{}{} kbConfig := map[string]interface{}{}
if kb.Instance != nil { if kb.Instance != nil {
if knowledgebase, ok := kb.Instance.(*kb.KnowledgeBase); ok && knowledgebase.Config != nil { if knowledgebase, ok := kb.Instance.(*kb.KnowledgeBase); ok && knowledgebase.Config != nil {
chunkings := []string{} // Use the current language setting for provider selection
if knowledgebase.Config.Chunkings != nil { currentLang := lang
for _, chunking := range knowledgebase.Config.Chunkings { if currentLang == "" {
chunkings = append(chunkings, chunking.ID) currentLang = "en" // Default to English
}
} }
embeddings := []string{} // Helper function to extract provider IDs from multi-language providers
if knowledgebase.Config.Embeddings != nil { extractProviderIDs := func(providerMap map[string][]*kbtypes.Provider) []string {
for _, embedding := range knowledgebase.Config.Embeddings { ids := []string{}
embeddings = append(embeddings, embedding.ID) 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{} var chunkings, embeddings, converters, extractors, fetchers []string
if knowledgebase.Config.Converters != nil { var searchers, rerankers, votes, weights, scores []string
for _, converter := range knowledgebase.Config.Converters {
converters = append(converters, converter.ID)
}
}
extractors := []string{} if knowledgebase.Providers != nil {
if knowledgebase.Config.Extractors != nil { chunkings = extractProviderIDs(knowledgebase.Providers.Chunkings)
for _, extractor := range knowledgebase.Config.Extractors { embeddings = extractProviderIDs(knowledgebase.Providers.Embeddings)
extractors = append(extractors, extractor.ID) converters = extractProviderIDs(knowledgebase.Providers.Converters)
} extractors = extractProviderIDs(knowledgebase.Providers.Extractors)
} fetchers = extractProviderIDs(knowledgebase.Providers.Fetchers)
searchers = extractProviderIDs(knowledgebase.Providers.Searchers)
fetchers := []string{} rerankers = extractProviderIDs(knowledgebase.Providers.Rerankers)
if knowledgebase.Config.Fetchers != nil { votes = extractProviderIDs(knowledgebase.Providers.Votes)
for _, fetcher := range knowledgebase.Config.Fetchers { weights = extractProviderIDs(knowledgebase.Providers.Weights)
fetchers = append(fetchers, fetcher.ID) scores = extractProviderIDs(knowledgebase.Providers.Scores)
}
}
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)
}
} }
kbConfig = map[string]interface{}{ kbConfig = map[string]interface{}{