Enhance knowledge base provider management with multi-language support
- Implemented loading of providers from directories, allowing for multi-language configurations. - Updated the KnowledgeBase struct to include a Providers field for managing provider configurations. - Refactored GetProviders and GetProvider methods to utilize the new multi-language provider system, improving localization support. - Added comprehensive test cases to validate provider loading and retrieval functionality across different languages.
This commit is contained in:
parent
0251bdc8d7
commit
4bdd921d4c
9 changed files with 728 additions and 314 deletions
115
kb/kb.go
115
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)
|
||||
}
|
||||
|
|
|
|||
127
kb/kb_test.go
127
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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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{}{
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue