diff --git a/plugins/golang-filter/mcp-server/servers/rag/orchestrator/orchestrator.go b/plugins/golang-filter/mcp-server/servers/rag/orchestrator/orchestrator.go deleted file mode 100644 index 5456eb0d..00000000 --- a/plugins/golang-filter/mcp-server/servers/rag/orchestrator/orchestrator.go +++ /dev/null @@ -1,198 +0,0 @@ -package orchestrator - -import ( - "context" - "strings" - "sync" - "time" - - "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/common/logger" - "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/config" - "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/crag" - "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/fusion" - "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/post" - pre_retrieve "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/pre-retrieve" - "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/retriever" - "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/schema" -) - -// Orchestrator wires the enhanced RAG pipeline stages. -type Orchestrator struct { - Cfg *config.Config - Retrievers []retriever.Retriever - Reranker post.Reranker - Compressor post.Compressor // Context compressor - Evaluator crag.Evaluator - WebSearcher *crag.WebSearcher - QueryRewriter *crag.QueryRewriter - Refiner *crag.KnowledgeRefiner - LLMProvider interface{} // Can be used for query rewriting and knowledge refinement - PreRetrieveProvider pre_retrieve.Provider // Pre-retrieve processor -} - -// Run executes the pipeline for a given query and returns final candidates. -func (o *Orchestrator) Run(ctx context.Context, query string) ([]schema.SearchResult, error) { - pc := o.Cfg.Pipeline - if pc == nil { - // No pipeline config; return empty to trigger fallback in caller. - return nil, nil - } - - // Pre-retrieve processing - subQueries := []string{query} - - if pc.EnablePre && o.PreRetrieveProvider != nil { - // Use the complete pre-retrieve provider (PreQRAG) - sessionID := "" // TODO: Extract from context or request if available - result, err := o.PreRetrieveProvider.Process(ctx, query, sessionID) - if err != nil { - // Log warning but continue with original query - logWarnf("Pre-retrieve processing failed: %v, using original query", err) - } else if result != nil { - // Extract queries from the plan nodes - if len(result.Plan.Nodes) > 0 { - subQueries = make([]string, 0, len(result.Plan.Nodes)) - for _, node := range result.Plan.Nodes { - // Use dense rewrite for vector retrieval - // For BM25/sparse retrieval, could use node.SparseRewrite - subQueries = append(subQueries, node.DenseRewrite) - } - - // Update query to aligned version for logging/later use - if result.AlignedQuery.Query != "" { - query = result.AlignedQuery.Query - } - } else { - // Fallback to aligned query if no plan nodes - if result.AlignedQuery.Query != "" { - query = result.AlignedQuery.Query - subQueries = []string{query} - } - } - } - } else if pc.EnablePre && o.Cfg.Pipeline.Pre != nil { - // Fallback to simple pre-processing if PreRetrieve not configured - // This is deprecated but kept for backward compatibility - if o.Cfg.Pipeline.Pre.Decompose.Enable { - // Simple decomposition: just use original query - subQueries = []string{query} - } - } - - // Hybrid retrieval - lists := make([][]schema.SearchResult, 0) - if pc.EnableHybrid { - for _, sq := range subQueries { - // Short timeout per sub-query; fan-out to retrievers in parallel - qctx, cancel := context.WithTimeout(ctx, 300*time.Millisecond) - var wg sync.WaitGroup - resCh := make(chan []schema.SearchResult, len(o.Retrievers)) - for _, r := range o.Retrievers { - rr := r - wg.Add(1) - go func() { - defer wg.Done() - if res, _ := rr.Search(qctx, sq, o.Cfg.RAG.TopK); len(res) > 0 { - resCh <- res - } - }() - } - wg.Wait() - close(resCh) - for res := range resCh { - lists = append(lists, res) - } - cancel() - } - } else { - // Minimal: use first retriever only - if len(o.Retrievers) > 0 { - res, _ := o.Retrievers[0].Search(ctx, query, o.Cfg.RAG.TopK) - if len(res) > 0 { - lists = append(lists, res) - } - } - } - - // Fuse - fused := fusion.RRFScore(lists, pc.RRFK) - - // Post-processing - if pc.EnablePost && o.Reranker != nil && o.Cfg.Pipeline.Post != nil && o.Cfg.Pipeline.Post.Rerank.Enable { - topN := o.Cfg.Pipeline.Post.Rerank.TopN - rr, _ := o.Reranker.Rerank(ctx, query, fused, topN) - fused = rr - } - - // Optional context compression - if pc.EnablePost && o.Cfg.Pipeline.Post != nil && o.Cfg.Pipeline.Post.Compress.Enable { - if o.Compressor != nil { - // Use advanced compressor with query awareness - compressed, err := o.Compressor.BatchCompress(ctx, fused, query) - if err != nil { - logWarnf("Compression failed: %v, using uncompressed results", err) - } else if len(compressed) > 0 { - fused = compressed - } - } else { - // Fallback to simple truncate compression (backward compatibility) - ratio := o.Cfg.Pipeline.Post.Compress.TargetRatio - for i := range fused { - fused[i].Document.Content = post.CompressText(fused[i].Document.Content, ratio) - } - } - } - - // CRAG - if pc.EnableCRAG && o.Evaluator != nil { - // Concatenate top-k contexts for quick evaluation - var b strings.Builder - limit := len(fused) - if limit > 5 { - limit = 5 - } - for i := 0; i < limit; i++ { - b.WriteString(fused[i].Document.Content) - b.WriteString("\n\n") - } - score, verdict, err := o.Evaluator.Evaluate(ctx, query, b.String()) - if err != nil { - // FailMode: closed -> bubble error, open -> keep fused - fm := "open" - if o.Cfg.Pipeline.CRAG != nil && o.Cfg.Pipeline.CRAG.FailMode != "" { - fm = o.Cfg.Pipeline.CRAG.FailMode - } - if fm == "closed" { - return nil, err - } - return fused, nil - } - _ = score // score could be logged/returned later - - // Build ActionContext for CRAG actions - actionCtx := &crag.ActionContext{ - Query: query, - Context: ctx, - WebSearcher: o.WebSearcher, - QueryRewriter: o.QueryRewriter, - Refiner: o.Refiner, - } - - // Execute appropriate corrective action based on verdict - switch verdict { - case crag.VerdictCorrect: - fused = crag.CorrectAction(actionCtx, fused) - case crag.VerdictIncorrect: - fused = crag.IncorrectAction(actionCtx) - case crag.VerdictAmbiguous: - fused = crag.AmbiguousAction(actionCtx, fused, nil) - } - } - - return fused, nil -} - -// Helper logging functions delegate to unified logger -func logWarnf(format string, args ...interface{}) { - logger.Warnf(format, args...) -} diff --git a/plugins/golang-filter/mcp-server/servers/rag/rag_client.go b/plugins/golang-filter/mcp-server/servers/rag/rag_client.go index f94c176c..6791ffac 100644 --- a/plugins/golang-filter/mcp-server/servers/rag/rag_client.go +++ b/plugins/golang-filter/mcp-server/servers/rag/rag_client.go @@ -20,7 +20,6 @@ import ( "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/gating" "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/llm" "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/metrics" - "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/orchestrator" "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/post" pre_retrieve "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/pre-retrieve" "github.com/alibaba/higress/plugins/golang-filter/mcp-server/servers/rag/profile" @@ -58,7 +57,17 @@ type RAGClient struct { cacheMode string indexVersion string cacheFusionVersion string - orch *orchestrator.Orchestrator + + // Post-processing components + compressor post.Compressor + + // CRAG components + webSearcher *crag.WebSearcher + queryRewriter *crag.QueryRewriter + refiner *crag.KnowledgeRefiner + + // Pre-retrieve component + preRetrieveProvider pre_retrieve.Provider } // NewRAGClient creates a new RAG client instance @@ -254,70 +263,59 @@ func NewRAGClient(config *config.Config) (*RAGClient, error) { } // Initialize reranker with support for multiple providers - var rr post.Reranker if ragclient.config.Pipeline.Post != nil && ragclient.config.Pipeline.Post.Rerank.Enable { rerankCfg := ragclient.config.Pipeline.Post.Rerank switch rerankCfg.Provider { case "llm": // Use LLM-based reranker if ragclient.llmProvider != nil { - rr = &post.LLMReranker{ + ragclient.reranker = &post.LLMReranker{ Provider: ragclient.llmProvider, Model: rerankCfg.Model, } } case "keyword": // Use keyword-based reranker - rr = &post.KeywordReranker{ + ragclient.reranker = &post.KeywordReranker{ MinKeywordLength: 3, BaseScoreWeight: 0.5, } case "model": // Use model-based reranker (BGE-reranker, Cohere rerank, etc.) - rr = &post.ModelReranker{ + ragclient.reranker = &post.ModelReranker{ Endpoint: rerankCfg.Endpoint, Model: rerankCfg.Model, APIKey: rerankCfg.APIKey, } default: // Default to HTTP reranker for backward compatibility - rr = post.NewHTTPReranker(rerankCfg.Endpoint) + ragclient.reranker = post.NewHTTPReranker(rerankCfg.Endpoint) } } - // Initialize CRAG components (for orchestrator) - var ev crag.Evaluator - var webSearcher *crag.WebSearcher - var queryRewriter *crag.QueryRewriter - var refiner *crag.KnowledgeRefiner - + // Initialize CRAG components if ragclient.config.Pipeline.CRAG != nil { cragCfg := ragclient.config.Pipeline.CRAG // Initialize evaluator (HTTP or LLM-based) if cragCfg.Evaluator.Provider == "http" && cragCfg.Evaluator.Endpoint != "" { - httpEval := &crag.HTTPEvaluator{ + ragclient.evaluator = &crag.HTTPEvaluator{ Endpoint: cragCfg.Evaluator.Endpoint, CorrectTh: cragCfg.Evaluator.Correct, IncorrectTh: cragCfg.Evaluator.Incorrect, } - ragclient.evaluator = httpEval - ev = httpEval } else if cragCfg.Evaluator.Provider == "llm" && ragclient.llmProvider != nil { - // Use LLM-based evaluator - llmEval := &crag.LLMEvaluator{ + ragclient.evaluator = &crag.LLMEvaluator{ Provider: ragclient.llmProvider, CorrectTh: cragCfg.Evaluator.Correct, IncorrectTh: cragCfg.Evaluator.Incorrect, } - ragclient.evaluator = llmEval - ev = llmEval } // Initialize web searcher from CRAG config or retriever config for _, rc := range ragclient.config.Pipeline.Retrievers { if rc.Type == "web" { - webSearcher = &crag.WebSearcher{ + ragclient.webSearcher = &crag.WebSearcher{ Provider: rc.Provider, Endpoint: rc.Params["endpoint"], APIKey: rc.Params["api_key"], @@ -326,19 +324,18 @@ func NewRAGClient(config *config.Config) (*RAGClient, error) { } } - // Initialize query rewriter if LLM available + // Initialize query rewriter and refiner if LLM available if ragclient.llmProvider != nil { - queryRewriter = &crag.QueryRewriter{ + ragclient.queryRewriter = &crag.QueryRewriter{ Provider: ragclient.llmProvider, } - refiner = &crag.KnowledgeRefiner{ + ragclient.refiner = &crag.KnowledgeRefiner{ Provider: ragclient.llmProvider, } } } // Initialize Compressor if enabled - var compressor post.Compressor if ragclient.config.Pipeline.Post != nil && ragclient.config.Pipeline.Post.Compress.Enable { compressCfg := ragclient.config.Pipeline.Post.Compress method := compressCfg.Method @@ -349,11 +346,10 @@ func NewRAGClient(config *config.Config) (*RAGClient, error) { if targetRatio == 0 { targetRatio = 0.7 // Default ratio } - compressor = post.NewCompressor(method, targetRatio, ragclient.llmProvider) + ragclient.compressor = post.NewCompressor(method, targetRatio, ragclient.llmProvider) } // Initialize Pre-Retrieve Provider if enabled - var preRetrieveProvider pre_retrieve.Provider if ragclient.config.Pipeline.EnablePre && ragclient.config.Pipeline.PreRetrieve != nil { preRetCfg := ragclient.config.Pipeline.PreRetrieve // Set LLM config if available @@ -366,22 +362,9 @@ func NewRAGClient(config *config.Config) (*RAGClient, error) { // Log warning but don't fail - pre-retrieve is optional fmt.Printf("[WARN] Failed to initialize pre-retrieve provider: %v\n", err) } else { - preRetrieveProvider = provider + ragclient.preRetrieveProvider = provider } } - - ragclient.orch = &orchestrator.Orchestrator{ - Cfg: ragclient.config, - Retrievers: retrievers, - Reranker: rr, - Compressor: compressor, - Evaluator: ev, - WebSearcher: webSearcher, - QueryRewriter: queryRewriter, - Refiner: refiner, - LLMProvider: ragclient.llmProvider, - PreRetrieveProvider: preRetrieveProvider, - } } return ragclient, nil } @@ -580,8 +563,49 @@ func (r *RAGClient) runEnhancedPipeline(ctx context.Context, query string) []sch } } - // Retrieval + // Pre-retrieve processing queries := []string{query} + originalQuery := query + if r.config.Pipeline != nil && r.config.Pipeline.EnablePre && r.preRetrieveProvider != nil { + sessionID := "" // TODO: Extract from context or request if available + result, err := r.preRetrieveProvider.Process(ctx, query, sessionID) + if err != nil { + api.LogWarnf("rag: pre-retrieve processing failed: %v, using original query", err) + } else if result != nil { + // Extract queries from the plan nodes + if len(result.Plan.Nodes) > 0 { + queries = make([]string, 0, len(result.Plan.Nodes)) + for _, node := range result.Plan.Nodes { + // Use dense rewrite for vector retrieval + // For BM25/sparse retrieval, could use node.SparseRewrite + if node.DenseRewrite != "" { + queries = append(queries, node.DenseRewrite) + } + } + if len(queries) == 0 { + queries = []string{query} + } + + // Update query to aligned version for logging/later use + if result.AlignedQuery.Query != "" { + originalQuery = result.AlignedQuery.Query + } + + if metricsRecord != nil { + metricsRecord.AddRetrievalPhase("pre_retrieve") + } + api.LogInfof("rag: pre-retrieve generated %d sub-queries from original query", len(queries)) + } else { + // Fallback to aligned query if no plan nodes + if result.AlignedQuery.Query != "" { + originalQuery = result.AlignedQuery.Query + queries = []string{originalQuery} + } + } + } + } + + // Retrieval results := r.retrievalProvider.Retrieve(ctx, queries, prof, metricsRecord) if metricsRecord != nil { @@ -601,7 +625,7 @@ func (r *RAGClient) runEnhancedPipeline(ctx context.Context, query string) []sch if topN <= 0 || topN > len(results) { topN = len(results) } - if reranked, err := r.reranker.Rerank(ctx, query, results, topN); err == nil && len(reranked) > 0 { + if reranked, err := r.reranker.Rerank(ctx, originalQuery, results, topN); err == nil && len(reranked) > 0 { results = reranked } if metricsRecord != nil { @@ -610,19 +634,30 @@ func (r *RAGClient) runEnhancedPipeline(ctx context.Context, query string) []sch } } - // Compression + // Compression with advanced compressor support if len(results) > 0 && r.config.Pipeline.EnablePost && r.config.Pipeline.Post != nil && r.config.Pipeline.Post.Compress.Enable { - ratio := r.config.Pipeline.Post.Compress.TargetRatio - for i := range results { - results[i].Document.Content = post.CompressText(results[i].Document.Content, ratio) + if r.compressor != nil { + // Use advanced compressor with query awareness + compressed, err := r.compressor.BatchCompress(ctx, results, originalQuery) + if err != nil { + api.LogWarnf("rag: compression failed: %v, using uncompressed results", err) + } else if len(compressed) > 0 { + results = compressed + } + } else { + // Fallback to simple truncate compression + ratio := r.config.Pipeline.Post.Compress.TargetRatio + for i := range results { + results[i].Document.Content = post.CompressText(results[i].Document.Content, ratio) + } } if metricsRecord != nil { metricsRecord.CompressEnabled = true } } - // CRAG evaluation + // CRAG evaluation with full action context if len(results) > 0 && r.config.Pipeline.EnableCRAG && r.evaluator != nil { var builder strings.Builder limit := len(results) @@ -633,17 +668,18 @@ func (r *RAGClient) runEnhancedPipeline(ctx context.Context, query string) []sch builder.WriteString(results[i].Document.Content) builder.WriteString("\n\n") } - _, verdict, err := r.evaluator.Evaluate(ctx, query, builder.String()) + _, verdict, err := r.evaluator.Evaluate(ctx, originalQuery, builder.String()) if err == nil { if r.feedbackManager != nil { r.feedbackManager.Record(prof.Name, verdict, 0) } // Build ActionContext for CRAG actions actionCtx := &crag.ActionContext{ - Query: query, - Context: ctx, - // WebSearcher, QueryRewriter, and Refiner would be available via orchestrator - // For direct RAGClient usage, they are optional + Query: originalQuery, + Context: ctx, + WebSearcher: r.webSearcher, + QueryRewriter: r.queryRewriter, + Refiner: r.refiner, } switch verdict { case crag.VerdictCorrect: