Merge pull request #1103 from trheyi/main

Refactor provider option resolution to require provider type
This commit is contained in:
Max 2025-08-13 16:08:55 +08:00 committed by GitHub
commit 6130907dec
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 17 additions and 21 deletions

View file

@ -16,7 +16,6 @@ func (p *Provider) GetOption(id string) (*ProviderOption, bool) {
return nil, false return nil, false
} }
// Find the option by id
for _, option := range p.Options { for _, option := range p.Options {
if option.Value == id { if option.Value == id {
return option, true return option, true

View file

@ -135,7 +135,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, locale string) (*kbtypes.ProviderOption, error) { func resolveProviderOption(config *ProviderConfig, providerType, 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")
} }
@ -144,6 +144,10 @@ func resolveProviderOption(config *ProviderConfig, locale string) (*kbtypes.Prov
return nil, fmt.Errorf("provider_id is required") return nil, fmt.Errorf("provider_id is required")
} }
if providerType == "" {
return nil, fmt.Errorf("provider_type is required")
}
// If Option is provided directly, use it // If Option is provided directly, use it
if config.Option != nil { if config.Option != nil {
return config.Option, nil return config.Option, nil
@ -159,14 +163,11 @@ func resolveProviderOption(config *ProviderConfig, locale string) (*kbtypes.Prov
locale = "en" locale = "en"
} }
// Find the provider using the new multi-language system // Find the provider using the specified provider type
var provider *kbtypes.Provider var provider *kbtypes.Provider
kbInstance := kb.Instance.(*kb.KnowledgeBase) kbInstance := kb.Instance.(*kb.KnowledgeBase)
// Check all provider types to find the matching provider // Get providers of the specific type
providerTypes := []string{"chunking", "embedding", "converter", "extractor", "fetcher"}
for _, providerType := range providerTypes {
providers := kbInstance.Providers.GetProviders(providerType, locale) providers := kbInstance.Providers.GetProviders(providerType, locale)
for _, p := range providers { for _, p := range providers {
if p.ID == config.ProviderID { if p.ID == config.ProviderID {
@ -174,10 +175,6 @@ func resolveProviderOption(config *ProviderConfig, locale string) (*kbtypes.Prov
break break
} }
} }
if provider != nil {
break
}
}
if provider == nil { if provider == nil {
return nil, fmt.Errorf("provider %s not found for locale %s", config.ProviderID, locale) return nil, fmt.Errorf("provider %s not found for locale %s", config.ProviderID, locale)
@ -232,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, locale) chunkingOption, err := resolveProviderOption(r.Chunking, "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)
} }
@ -251,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, locale) embeddingOption, err := resolveProviderOption(r.Embedding, "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)
} }
@ -264,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, locale) extractionOption, err := resolveProviderOption(r.Extraction, "extractor", 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)
} }
@ -277,7 +274,7 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp
} }
if r.Fetcher != nil { if r.Fetcher != nil {
fetcherOption, err := resolveProviderOption(r.Fetcher, locale) fetcherOption, err := resolveProviderOption(r.Fetcher, "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)
} }
@ -292,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, locale) converterOption, err := resolveProviderOption(r.Converter, "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)
} }
@ -315,7 +312,7 @@ func (r *BaseUpsertRequest) ToUpsertOptions(fileInfo ...string) (*types.UpsertOp
ProviderID: converterID, ProviderID: converterID,
} }
converterOption, err := resolveProviderOption(converterConfig, locale) converterOption, err := resolveProviderOption(converterConfig, "converter", 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)
} }