diff --git a/api/apps/restful_apis/bot_api.py b/api/apps/restful_apis/bot_api.py index 2c1938357b..b52f7f9217 100644 --- a/api/apps/restful_apis/bot_api.py +++ b/api/apps/restful_apis/bot_api.py @@ -44,6 +44,7 @@ from rag.prompts.template import load_prompt from rag.prompts.generator import cross_languages, keyword_extraction from common.constants import RetCode, LLMType, StatusEnum from common import settings +from rag.utils.web_search_conn import has_web_search_provider from api.utils.reference_metadata_utils import ( enrich_chunks_with_document_metadata, resolve_reference_metadata_preferences, @@ -141,12 +142,14 @@ async def chatbots_inputs(dialog_id, tenant_id=None): request_session_id, ) return get_error_data_result(message="Authentication error: no access to this chatbot!") + has_web_search = has_web_search_provider(dialog.prompt_config) return get_result( data={ "title": dialog.name, "avatar": dialog.icon, "prologue": dialog.prompt_config.get("prologue", ""), - "has_tavily_key": bool(dialog.prompt_config.get("tavily_api_key", "").strip()), + "has_tavily_key": has_web_search, + "has_web_search_provider": has_web_search, "llm_id": dialog.llm_id or "", } ) diff --git a/api/db/services/dialog_service.py b/api/db/services/dialog_service.py index 47e65b2a8f..bde5d2858d 100644 --- a/api/db/services/dialog_service.py +++ b/api/db/services/dialog_service.py @@ -50,7 +50,7 @@ from rag.app.tag import label_question from rag.nlp.search import index_name from rag.prompts.generator import chunks_format, citation_prompt, cross_languages, full_question, kb_prompt, keyword_extraction, message_fit_in, PROMPT_JINJA_ENV, ASK_SUMMARY from common.token_utils import num_tokens_from_string -from rag.utils.tavily_conn import Tavily +from rag.utils.web_search_conn import create_web_search_provider, has_web_search_provider from rag.utils.tts_cache import synthesize_with_cache from common.string_utils import remove_redundant_spaces from common import settings @@ -124,7 +124,7 @@ def _normalize_internet_flag(value): def _should_use_web_search(prompt_config, internet=None): - if not prompt_config.get("tavily_api_key"): + if not has_web_search_provider(prompt_config): return False normalized = _normalize_internet_flag(internet) return normalized is True @@ -577,7 +577,7 @@ async def async_chat(dialog, messages, stream=True, **kwargs): assert messages[-1]["role"] == "user", "The last content of this conversation is not from user." session_id = kwargs.get("session_id") use_web_search = _should_use_web_search(dialog.prompt_config, kwargs.get("internet")) - logging.debug("web_search kb=%s tavily=%s internet=%r enabled=%s", bool(dialog.kb_ids), bool(dialog.prompt_config.get("tavily_api_key")), kwargs.get("internet"), use_web_search) + logging.debug("web_search kb=%s configured=%s internet=%r enabled=%s", bool(dialog.kb_ids), has_web_search_provider(dialog.prompt_config), kwargs.get("internet"), use_web_search) if not dialog.kb_ids and not use_web_search: async for ans in async_chat_solo(dialog, messages, stream, session_id=session_id): yield ans @@ -782,10 +782,10 @@ async def async_chat(dialog, messages, stream=True, **kwargs): kbinfos["chunks"] = cks kbinfos["chunks"] = retriever.retrieval_by_children(kbinfos["chunks"], tenant_ids) if use_web_search: - tav = Tavily(prompt_config["tavily_api_key"]) - tav_res = tav.retrieve_chunks(" ".join(questions)) - kbinfos["chunks"].extend(tav_res["chunks"]) - kbinfos["doc_aggs"].extend(tav_res["doc_aggs"]) + web_search = create_web_search_provider(prompt_config) + web_res = web_search.retrieve_chunks(" ".join(questions)) + kbinfos["chunks"].extend(web_res["chunks"]) + kbinfos["doc_aggs"].extend(web_res["doc_aggs"]) if prompt_config.get("use_kg"): default_chat_model = get_tenant_default_model_by_type(dialog.tenant_id, LLMType.CHAT) ck = await settings.kg_retriever.retrieval( @@ -1879,7 +1879,7 @@ async def rag_agent(dialog, messages, stream=True, **kwargs): return kbs, embd_mdl, rerank_mdl, chat_mdl, tts_mdl = get_models(dialog) use_web_search = _should_use_web_search(prompt_config, kwargs.get("internet")) - logging.debug("web_search kb=%s tavily=%s internet=%r enabled=%s", bool(dialog.kb_ids), bool(prompt_config.get("tavily_api_key")), kwargs.get("internet"), use_web_search) + logging.debug("web_search kb=%s configured=%s internet=%r enabled=%s", bool(dialog.kb_ids), has_web_search_provider(prompt_config), kwargs.get("internet"), use_web_search) tenant_ids = list(set([kb.tenant_id for kb in kbs])) # "reasoning" arrives as "1".."4" mapping to the ordered THINKING_MODES # (low, medium, high, ultra); fall back to "medium" on anything else. @@ -1917,7 +1917,7 @@ async def rag_agent(dialog, messages, stream=True, **kwargs): chat_mdl, embed_mdl=embd_mdl, kb_ids=dialog.kb_ids, - tav=Tavily(prompt_config.get("tavily_api_key")) if use_web_search else None, + web_search=create_web_search_provider(prompt_config) if use_web_search else None, meta_data_filter=dialog.meta_data_filter, doc_scope=doc_scope, do_refer=False, diff --git a/docs/references/http_api_reference.md b/docs/references/http_api_reference.md index 3b3a9c8fdd..1700543b3f 100644 --- a/docs/references/http_api_reference.md +++ b/docs/references/http_api_reference.md @@ -3032,7 +3032,9 @@ curl --request POST \ - `"use_kg"`: `boolean` - `"reasoning"`: `boolean` - `"cross_languages"`: `list[string]` + - `"web_search_provider"`: `string` The web search service to use. Supported values are `"tavily"` and `"querit"`. Defaults to `"tavily"` when omitted. - `"tavily_api_key"`: `string` + - `"querit_api_key"`: `string` The Querit API key. Set `web_search_provider` to `"querit"` when using this field. - `"toc_enhance"`: `boolean` - `"similarity_threshold"`: (*Body parameter*), `float` - `"vector_similarity_weight"`: (*Body parameter*), `float` diff --git a/internal/handler/bot.go b/internal/handler/bot.go index 51c4236cdf..b49bcb94b2 100644 --- a/internal/handler/bot.go +++ b/internal/handler/bot.go @@ -39,7 +39,7 @@ type BotHandler struct { // is interface-typed so the test suite can inject a stub. type botService interface { ChatbotInfo(ctx context.Context, tenantID, dialogID string) ( - title, avatar, prologue, llmID string, hasTavilyKey bool, ec common.ErrorCode, err error) + title, avatar, prologue, llmID string, hasWebSearch bool, ec common.ErrorCode, err error) AgentbotInputs(ctx context.Context, tenantID, agentID string) ( title, avatar, prologue, mode string, inputs map[string]any, ec common.ErrorCode, err error) @@ -58,7 +58,7 @@ func NewBotHandler(svc *service.BotService) *BotHandler { // ChatbotInfo GET /api/v1/chatbots//info // // Mirrors python bot_api.py:126-154. Returns the public metadata of -// a chatbot dialog (title, avatar, prologue, tavily key flag, llm_id). +// a chatbot dialog (title, avatar, prologue, web search flag, llm_id). func (h *BotHandler) ChatbotInfo(c *gin.Context) { user, code, msg := GetUser(c) if code != common.CodeSuccess { @@ -70,18 +70,19 @@ func (h *BotHandler) ChatbotInfo(c *gin.Context) { common.ResponseWithCodeData(c, common.CodeArgumentError, nil, "`dialog_id` is required.") return } - title, avatar, prologue, llmID, hasTavily, ec, err := h.botService.ChatbotInfo( + title, avatar, prologue, llmID, hasWebSearch, ec, err := h.botService.ChatbotInfo( c.Request.Context(), user.ID, dialogID) if err != nil { common.ResponseWithCodeData(c, ec, nil, err.Error()) return } common.SuccessWithData(c, gin.H{ - "title": title, - "avatar": avatar, - "prologue": prologue, - "has_tavily_key": hasTavily, - "llm_id": llmID, + "title": title, + "avatar": avatar, + "prologue": prologue, + "has_tavily_key": hasWebSearch, + "has_web_search_provider": hasWebSearch, + "llm_id": llmID, }, "success") } diff --git a/internal/handler/bot_test.go b/internal/handler/bot_test.go index ff2b77a990..073998e12d 100644 --- a/internal/handler/bot_test.go +++ b/internal/handler/bot_test.go @@ -184,6 +184,9 @@ func TestChatbotInfo_HasTavilyKey(t *testing.T) { if resp.Data["has_tavily_key"] != true { t.Errorf("has_tavily_key = %v, want true", resp.Data["has_tavily_key"]) } + if resp.Data["has_web_search_provider"] != true { + t.Errorf("has_web_search_provider = %v, want true", resp.Data["has_web_search_provider"]) + } } // TestChatbotInfo_ForeignTenant covers criterion 15. diff --git a/internal/service/bot.go b/internal/service/bot.go index e1bbaf6452..3c8adf7d53 100644 --- a/internal/service/bot.go +++ b/internal/service/bot.go @@ -28,7 +28,6 @@ import ( "errors" "fmt" "hash/fnv" - "strings" "sync" "ragflow/internal/agent/canvas" @@ -83,7 +82,7 @@ func NewBotService(agentSvc *AgentService, llmSvc *LLMService) *BotService { // it (TenantID match), and Status must equal common.StatusDialogValid // (the python StatusEnum.VALID.value). func (s *BotService) ChatbotInfo(ctx context.Context, tenantID, dialogID string) ( - title, avatar, prologue, llmID string, hasTavilyKey bool, ec common.ErrorCode, err error, + title, avatar, prologue, llmID string, hasWebSearch bool, ec common.ErrorCode, err error, ) { dialog, err := s.chatDAO.GetDialogByID(ctx, dao.DB, dialogID) if err != nil { @@ -97,14 +96,13 @@ func (s *BotService) ChatbotInfo(ctx context.Context, tenantID, dialogID string) pc := dialog.PromptConfig // Defensive lookups mirroring python's // dialog.prompt_config.get("prologue", "") and - // dialog.prompt_config.get("tavily_api_key", "").strip() + // resolveWebSearchProvider(dialog.prompt_config) != nil // semantics. A hard type assertion here would panic on a missing // or non-string prologue field — this endpoint is public over // persisted JSON config and the schema is not guaranteed. prologue = stringFromMap(pc, "prologue") - tk := stringFromMap(pc, "tavily_api_key") return botDerefStr(dialog.Name), botDerefStr(dialog.Icon), prologue, - dialog.LLMID, strings.TrimSpace(tk) != "", common.CodeSuccess, nil + dialog.LLMID, resolveWebSearchProvider(pc) != nil, common.CodeSuccess, nil } // AgentbotInputs returns the public metadata of an agentbot canvas. diff --git a/internal/service/chat_pipeline.go b/internal/service/chat_pipeline.go index 2ac0c6bc83..88fb8223b6 100644 --- a/internal/service/chat_pipeline.go +++ b/internal/service/chat_pipeline.go @@ -125,12 +125,12 @@ type AsyncChatResult struct { // │ │ // │ reasoning=true? │ // │ YES → DeepResearcher (recursive, maxDepth=3) │ -// │ each layer: KB → Web(Tavily) → KG(use_kg) │ +// │ each layer: KB → Web search → KG(use_kg) │ // │ → sufficiencyCheck → multiQueriesGen → recurse│ // │ NO → Standard vector retrieval │ // │ vector/hybrid search → rerank → │ // │ TOC enhance → child chunk retrieval → │ -// │ Tavily web search → KG retrieval (prepend) │ +// │ Web search → KG retrieval (prepend) │ // │ │ // │ enrichChunksWithMetadata (doc metadata) │ // │ kbPrompt (build knowledge blocks) │ @@ -184,7 +184,7 @@ func (s *ChatPipelineService) AsyncChat( if useWebSearch { common.Debug("web_search", zap.Bool("kb", hasKBs), - zap.Bool("tavily", chat.PromptConfig != nil && chat.PromptConfig["tavily_api_key"] != "" && chat.PromptConfig["tavily_api_key"] != nil), + zap.Bool("configured", resolveWebSearchProvider(chat.PromptConfig) != nil), zap.Any("internet", kwargs["internet"]), zap.Bool("enabled", useWebSearch)) } @@ -612,7 +612,7 @@ func (s *ChatPipelineService) AsyncChat( // b) Otherwise: standard retrieval, then: // - TOC enhancement (if toc_enhance is enabled). // - Child chunk retrieval. - // - Tavily web search (if internet is enabled). + // - Web search provider (if internet is enabled). // - Knowledge graph retrieval (if use_kg is enabled). // Populates kbinfos (chunks + doc_aggs) and knowledges. // When false, the entire block is skipped. @@ -772,21 +772,21 @@ func (s *ChatPipelineService) AsyncChat( kbinfos["chunks"] = nlp.RetrievalByChildren(existingChunks, kbTenantIDStrings(kbs), engine.Get(), ctx) } - // Web search via Tavily + // Web search if s.shouldUseWebSearch(chat, kwargs["internet"]) { - tavilyKey, _ := chat.PromptConfig["tavily_api_key"].(string) - tavResult, tavErr := s.tavilyRetrieve(ctx, tavilyKey, searchQuestion) - if tavErr != nil { - common.Warn("Tavily web search failed", zap.Error(tavErr)) + provider := resolveWebSearchProvider(chat.PromptConfig) + webResult, webErr := s.retrieveWebSearch(ctx, provider, searchQuestion) + if webErr != nil { + common.Warn("Web search failed", zap.Error(webErr)) } else { // Extend chunks and doc_aggs with web search results. if existingChunks, ok := kbinfos["chunks"].([]map[string]interface{}); ok { - if newChunks, ok := tavResult["chunks"].([]map[string]interface{}); ok { + if newChunks, ok := webResult["chunks"].([]map[string]interface{}); ok { kbinfos["chunks"] = append(existingChunks, newChunks...) } } if existingAggs, ok := kbinfos["doc_aggs"].([]interface{}); ok { - if newAggs, ok := tavResult["doc_aggs"].([]interface{}); ok { + if newAggs, ok := webResult["doc_aggs"].([]interface{}); ok { kbinfos["doc_aggs"] = append(existingAggs, newAggs...) } } @@ -1755,7 +1755,7 @@ func normalizeInternetFlag(v interface{}) *bool { // shouldUseWebSearch returns true if web search should be enabled. // Mirrors Python's _should_use_web_search (dialog_service.py:122-126): -// Tavily key must be present on chat.PromptConfig AND the internet +// A web search provider must be configured on chat.PromptConfig AND the internet // flag must normalize to explicit true. // // The second parameter takes the raw internet value (typically @@ -1765,8 +1765,7 @@ func (s *ChatPipelineService) shouldUseWebSearch(chat *entity.Chat, internet int if chat.PromptConfig == nil { return false } - tavilyKey, _ := chat.PromptConfig["tavily_api_key"].(string) - if tavilyKey == "" { + if resolveWebSearchProvider(chat.PromptConfig) == nil { return false } normalized := normalizeInternetFlag(internet) diff --git a/internal/service/deep_researcher.go b/internal/service/deep_researcher.go index 38897e670c..4593bd793c 100644 --- a/internal/service/deep_researcher.go +++ b/internal/service/deep_researcher.go @@ -139,7 +139,7 @@ type DeepResearcher struct { PromptConfig map[string]interface{} KBRetrieve KBRetrieveFunc InternetEnabled bool - TavilyAPIKey string + WebSearch *webSearchProviderConfig // Fields needed for KG retrieval (mirrors async_chat.go usage). DocEngine engine.DocEngine @@ -168,7 +168,7 @@ func NewDeepResearcher( PromptConfig: promptConfig, KBRetrieve: kbRetrieve, InternetEnabled: internetEnabled, - TavilyAPIKey: mapStringValue(promptConfig, "tavily_api_key"), + WebSearch: resolveWebSearchProvider(promptConfig), DocEngine: docEngine, KbIDs: kbIDs, TenantIDs: tenantIDs, @@ -367,17 +367,17 @@ func (dr *DeepResearcher) retrieveInformation(ctx context.Context, query string) } } - // 2. Web retrieval (Tavily) - if dr.InternetEnabled && dr.TavilyAPIKey != "" { - tavRes, err := dr.tavilyRetrieve(ctx, query) + // 2. Web retrieval + if dr.InternetEnabled && dr.WebSearch != nil { + webRes, err := dr.retrieveWebSearch(ctx, dr.WebSearch, query) if err != nil { common.Warn("DeepResearcher: web retrieval error", zap.Error(err)) - } else if tavRes != nil { - if chunks, ok := tavRes["chunks"].([]map[string]interface{}); ok { + } else if webRes != nil { + if chunks, ok := webRes["chunks"].([]map[string]interface{}); ok { existing, _ := kbinfos["chunks"].([]map[string]interface{}) kbinfos["chunks"] = append(existing, chunks...) } - if aggs, ok := tavRes["doc_aggs"].([]interface{}); ok { + if aggs, ok := webRes["doc_aggs"].([]interface{}); ok { existing, _ := kbinfos["doc_aggs"].([]interface{}) kbinfos["doc_aggs"] = append(existing, aggs...) } @@ -410,10 +410,10 @@ func (dr *DeepResearcher) retrieveInformation(ctx context.Context, query string) } // tavilyRetrieve calls the Tavily Search API. -func (dr *DeepResearcher) tavilyRetrieve(ctx context.Context, query string) (map[string]interface{}, error) { +func (dr *DeepResearcher) tavilyRetrieve(ctx context.Context, apiKey, query string) (map[string]interface{}, error) { reqBody := map[string]interface{}{ "query": query, - "api_key": dr.TavilyAPIKey, + "api_key": apiKey, "search_depth": "advanced", "max_results": 6, } @@ -836,16 +836,6 @@ func getMapString(m map[string]interface{}, keys ...string) string { return "" } -// mapStringValue extracts a string value from a map by key. -func mapStringValue(m map[string]interface{}, key string) string { - if v, ok := m[key]; ok { - if s, ok := v.(string); ok { - return s - } - } - return "" -} - // chunksFromKBInfos extracts chunks list from kbinfos for counting. func chunksFromKBInfos(kbinfos map[string]interface{}) []map[string]interface{} { if ch, ok := kbinfos["chunks"].([]map[string]interface{}); ok { diff --git a/internal/service/web_search_provider.go b/internal/service/web_search_provider.go new file mode 100644 index 0000000000..85114856f9 --- /dev/null +++ b/internal/service/web_search_provider.go @@ -0,0 +1,246 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package service + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" +) + +const ( + webSearchProviderTavily = "tavily" + webSearchProviderQuerit = "querit" + queritWebSearchEndpoint = "https://api.querit.ai/v1/search" +) + +var queritWebSearchHTTPClient = &http.Client{Timeout: 30 * time.Second} + +type webSearchProviderConfig struct { + Provider string + APIKey string +} + +func resolveWebSearchProvider(promptConfig map[string]interface{}) *webSearchProviderConfig { + if promptConfig == nil { + return nil + } + + provider := webSearchProviderTavily + if configuredProvider, exists := promptConfig["web_search_provider"]; exists { + var ok bool + provider, ok = configuredProvider.(string) + if !ok { + return nil + } + } + + apiKeyField := "" + switch provider { + case webSearchProviderTavily: + apiKeyField = "tavily_api_key" + case webSearchProviderQuerit: + apiKeyField = "querit_api_key" + default: + return nil + } + + apiKey, _ := promptConfig[apiKeyField].(string) + apiKey = strings.TrimSpace(apiKey) + if apiKey == "" { + return nil + } + return &webSearchProviderConfig{ + Provider: provider, + APIKey: apiKey, + } +} + +func (s *ChatPipelineService) retrieveWebSearch( + ctx context.Context, + provider *webSearchProviderConfig, + question string, +) (map[string]interface{}, error) { + if provider == nil { + return nil, fmt.Errorf("web search provider is not configured") + } + switch provider.Provider { + case webSearchProviderTavily: + return s.tavilyRetrieve(ctx, provider.APIKey, question) + case webSearchProviderQuerit: + return retrieveQueritWebSearch( + ctx, + queritWebSearchHTTPClient, + queritWebSearchEndpoint, + provider.APIKey, + question, + ) + default: + return nil, fmt.Errorf("unsupported web search provider %q", provider.Provider) + } +} + +func (dr *DeepResearcher) retrieveWebSearch( + ctx context.Context, + provider *webSearchProviderConfig, + query string, +) (map[string]interface{}, error) { + if provider == nil { + return nil, fmt.Errorf("web search provider is not configured") + } + switch provider.Provider { + case webSearchProviderTavily: + return dr.tavilyRetrieve(ctx, provider.APIKey, query) + case webSearchProviderQuerit: + return retrieveQueritWebSearch( + ctx, + queritWebSearchHTTPClient, + queritWebSearchEndpoint, + provider.APIKey, + query, + ) + default: + return nil, fmt.Errorf("unsupported web search provider %q", provider.Provider) + } +} + +type queritWebSearchResult struct { + Title string `json:"title"` + URL string `json:"url"` + Snippet string `json:"snippet"` +} + +func retrieveQueritWebSearch( + ctx context.Context, + client *http.Client, + endpoint string, + apiKey string, + query string, +) (map[string]interface{}, error) { + requestBody, err := json.Marshal(map[string]interface{}{ + "query": query, + "count": 6, + "chunksPerDoc": 1, + }) + if err != nil { + return nil, fmt.Errorf("querit: marshal request: %w", err) + } + + request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(requestBody)) + if err != nil { + return nil, fmt.Errorf("querit: new request: %w", err) + } + request.Header.Set("Accept", "application/json") + request.Header.Set("Authorization", "Bearer "+apiKey) + request.Header.Set("Content-Type", "application/json") + + response, err := client.Do(request) + if err != nil { + return nil, fmt.Errorf("querit: do request: %w", err) + } + defer response.Body.Close() + + if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices { + return nil, fmt.Errorf("querit: status %d", response.StatusCode) + } + + responseBody, err := io.ReadAll(response.Body) + if err != nil { + return nil, fmt.Errorf("querit: read response: %w", err) + } + results, err := decodeQueritWebSearchResults(responseBody) + if err != nil { + return nil, err + } + + chunks := make([]map[string]interface{}, 0, len(results)) + docAggs := make([]interface{}, 0, len(results)) + for _, result := range results { + if result.Snippet == "" { + continue + } + chunkID := "querit-" + result.URL + chunks = append(chunks, map[string]interface{}{ + "chunk_id": chunkID, + "content_ltks": tokenizeText(result.Snippet), + "content_with_weight": result.Snippet, + "doc_id": chunkID, + "docnm_kwd": result.Title, + "kb_id": []interface{}{}, + "important_kwd": []interface{}{}, + "image_id": "", + "similarity": float64(1), + "vector_similarity": float64(1), + "term_similarity": float64(0), + "vector": []float64{}, + "positions": []interface{}{}, + "url": result.URL, + }) + docAggs = append(docAggs, map[string]interface{}{ + "doc_name": result.Title, + "doc_id": chunkID, + "count": 1, + "url": result.URL, + }) + } + + return map[string]interface{}{ + "chunks": chunks, + "doc_aggs": docAggs, + }, nil +} + +func decodeQueritWebSearchResults(responseBody []byte) ([]queritWebSearchResult, error) { + var envelope map[string]json.RawMessage + if err := json.Unmarshal(responseBody, &envelope); err != nil { + return nil, fmt.Errorf("querit: decode response: %w", err) + } + if envelope == nil { + return nil, fmt.Errorf("querit: response must be an object") + } + + resultsValue, exists := envelope["results"] + if !exists { + return []queritWebSearchResult{}, nil + } + if strings.TrimSpace(string(resultsValue)) == "null" { + return nil, fmt.Errorf("querit: response field results must be an object") + } + + var resultsContainer map[string]json.RawMessage + if err := json.Unmarshal(resultsValue, &resultsContainer); err != nil { + return nil, fmt.Errorf("querit: response field results must be an object: %w", err) + } + resultValue, exists := resultsContainer["result"] + if !exists { + return []queritWebSearchResult{}, nil + } + if strings.TrimSpace(string(resultValue)) == "null" { + return nil, fmt.Errorf("querit: response field results.result must be an array") + } + + var results []queritWebSearchResult + if err := json.Unmarshal(resultValue, &results); err != nil { + return nil, fmt.Errorf("querit: response field results.result must be an array: %w", err) + } + return results, nil +} diff --git a/internal/service/web_search_provider_test.go b/internal/service/web_search_provider_test.go new file mode 100644 index 0000000000..d503d48b0d --- /dev/null +++ b/internal/service/web_search_provider_test.go @@ -0,0 +1,222 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package service + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" +) + +func TestResolveWebSearchProviderUsesExistingTavilyConfig(t *testing.T) { + provider := resolveWebSearchProvider(map[string]interface{}{ + "tavily_api_key": "tvly-test", + }) + + if provider == nil { + t.Fatal("provider is nil") + } + if provider.Provider != webSearchProviderTavily { + t.Fatalf("provider = %q, want %q", provider.Provider, webSearchProviderTavily) + } + if provider.APIKey != "tvly-test" { + t.Fatalf("api key = %q, want %q", provider.APIKey, "tvly-test") + } +} + +func TestResolveWebSearchProviderReturnsNilWithoutTavilyKey(t *testing.T) { + cases := []struct { + name string + config map[string]interface{} + }{ + {name: "nil config", config: nil}, + {name: "empty config", config: map[string]interface{}{}}, + {name: "empty key", config: map[string]interface{}{"tavily_api_key": ""}}, + {name: "whitespace key", config: map[string]interface{}{"tavily_api_key": " "}}, + {name: "non-string key", config: map[string]interface{}{"tavily_api_key": 1}}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if provider := resolveWebSearchProvider(tc.config); provider != nil { + t.Fatalf("provider = %+v, want nil", provider) + } + }) + } +} + +func TestResolveWebSearchProviderUsesSelectedQueritConfig(t *testing.T) { + provider := resolveWebSearchProvider(map[string]interface{}{ + "web_search_provider": "querit", + "querit_api_key": "querit-test", + "tavily_api_key": "tvly-test", + }) + + if provider == nil { + t.Fatal("provider is nil") + } + if provider.Provider != webSearchProviderQuerit { + t.Fatalf("provider = %q, want %q", provider.Provider, webSearchProviderQuerit) + } + if provider.APIKey != "querit-test" { + t.Fatalf("api key = %q, want %q", provider.APIKey, "querit-test") + } +} + +func TestResolveWebSearchProviderTrimsSelectedKey(t *testing.T) { + provider := resolveWebSearchProvider(map[string]interface{}{ + "web_search_provider": "querit", + "querit_api_key": " querit-test ", + }) + + if provider == nil { + t.Fatal("provider is nil") + } + if provider.APIKey != "querit-test" { + t.Fatalf("api key = %q, want %q", provider.APIKey, "querit-test") + } +} + +func TestResolveWebSearchProviderRequiresKeyForSelectedProvider(t *testing.T) { + cases := []struct { + name string + config map[string]interface{} + }{ + {name: "tavily", config: map[string]interface{}{"web_search_provider": "tavily"}}, + {name: "querit", config: map[string]interface{}{"web_search_provider": "querit"}}, + { + name: "querit whitespace key", + config: map[string]interface{}{ + "web_search_provider": "querit", + "querit_api_key": " ", + }, + }, + { + name: "querit does not fall back to tavily", + config: map[string]interface{}{ + "web_search_provider": "querit", + "tavily_api_key": "tvly-test", + }, + }, + { + name: "unsupported provider", + config: map[string]interface{}{ + "web_search_provider": "unsupported", + "querit_api_key": "querit-test", + "tavily_api_key": "tvly-test", + }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if provider := resolveWebSearchProvider(tc.config); provider != nil { + t.Fatalf("provider = %+v, want nil", provider) + } + }) + } +} + +func TestRetrieveQueritWebSearchUsesChatDefaultsAndReturnsReferenceShape(t *testing.T) { + var requestBody map[string]interface{} + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + if got := request.Header.Get("Authorization"); got != "Bearer querit-test" { + t.Errorf("Authorization = %q, want %q", got, "Bearer querit-test") + } + if err := json.NewDecoder(request.Body).Decode(&requestBody); err != nil { + t.Errorf("decode request: %v", err) + return + } + response.Header().Set("Content-Type", "application/json") + _, _ = response.Write([]byte(`{ + "results": { + "result": [{ + "title": "RAGFlow", + "url": "https://example.com/ragflow", + "snippet": "RAGFlow is an open-source RAG engine." + }] + } + }`)) + })) + defer server.Close() + + result, err := retrieveQueritWebSearch( + context.Background(), + server.Client(), + server.URL, + "querit-test", + "What is RAGFlow?", + ) + if err != nil { + t.Fatalf("retrieve Querit web search: %v", err) + } + + if requestBody["query"] != "What is RAGFlow?" { + t.Fatalf("query = %#v, want %q", requestBody["query"], "What is RAGFlow?") + } + if requestBody["count"] != float64(6) { + t.Fatalf("count = %#v, want 6", requestBody["count"]) + } + if requestBody["chunksPerDoc"] != float64(1) { + t.Fatalf("chunksPerDoc = %#v, want 1", requestBody["chunksPerDoc"]) + } + + chunks, ok := result["chunks"].([]map[string]interface{}) + if !ok || len(chunks) != 1 { + t.Fatalf("chunks = %#v, want one chunk", result["chunks"]) + } + if chunks[0]["content_with_weight"] != "RAGFlow is an open-source RAG engine." { + t.Fatalf("content = %#v", chunks[0]["content_with_weight"]) + } + if chunks[0]["docnm_kwd"] != "RAGFlow" { + t.Fatalf("title = %#v", chunks[0]["docnm_kwd"]) + } + if chunks[0]["url"] != "https://example.com/ragflow" { + t.Fatalf("url = %#v", chunks[0]["url"]) + } + if chunks[0]["similarity"] != float64(1) { + t.Fatalf("similarity = %#v, want 1", chunks[0]["similarity"]) + } + + aggs, ok := result["doc_aggs"].([]interface{}) + if !ok || len(aggs) != 1 { + t.Fatalf("doc_aggs = %#v, want one aggregate", result["doc_aggs"]) + } +} + +func TestDecodeQueritWebSearchResultsRejectsMalformedContainers(t *testing.T) { + cases := []struct { + name string + body string + }{ + {name: "null response", body: `null`}, + {name: "null results", body: `{"results":null}`}, + {name: "array results", body: `{"results":[]}`}, + {name: "null result list", body: `{"results":{"result":null}}`}, + {name: "object result list", body: `{"results":{"result":{}}}`}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if _, err := decodeQueritWebSearchResults([]byte(tc.body)); err == nil { + t.Fatal("error is nil") + } + }) + } +} diff --git a/rag/advanced_rag/agentic_rag.py b/rag/advanced_rag/agentic_rag.py index 04ea22dde5..2c7a0e61c1 100644 --- a/rag/advanced_rag/agentic_rag.py +++ b/rag/advanced_rag/agentic_rag.py @@ -55,7 +55,7 @@ from rag.prompts.generator import ( sufficiency_select, ) from api.db.db_models import Document, Knowledgebase -from rag.utils.tavily_conn import Tavily +from rag.utils.web_search_conn import WebSearchProvider # Tokens held back from the model's context when fitting retrieved evidence @@ -75,7 +75,7 @@ class RAGTools: embed_mdl: LLMBundle | None = None, kb_ids: List[str] | None = None, kbs: list[Knowledgebase] | None = None, - tav: Tavily | None = None, + web_search: WebSearchProvider | None = None, meta_data_filter: dict | None = None, doc_scope: List[str] | None = None, user_defined_prompts: dict | None = None, @@ -106,7 +106,7 @@ class RAGTools: for kb in kbs: _exclude_sql_kb(kb) - self.tav = tav + self.web_search = web_search self.meta_data_filter = meta_data_filter self.doc_scope = list(dict.fromkeys(doc_scope)) if doc_scope is not None else None self.user_defined_prompts = user_defined_prompts or {} @@ -137,7 +137,7 @@ class RAGTools: return bool(self.sql_kbs and self.field_map) def has_web(self) -> bool: - return self.tav is not None + return self.web_search is not None def has_llm(self) -> bool: return self.chat_mdl is not None @@ -423,15 +423,15 @@ class RAGTools: return {"chunks": kbinfos.get("chunks", []), "doc_aggs": kbinfos.get("doc_aggs", [])} async def web_retrieve(self, query: str) -> dict[str, list]: - """Retrieve chunks from the public web (Tavily). Raw kbinfos shape.""" - if self.tav is None: + """Retrieve chunks from the public web. Raw kbinfos shape.""" + if self.web_search is None: return {"chunks": [], "doc_aggs": []} try: - tav_res = await thread_pool_exec(self.tav.retrieve_chunks, query) + web_res = await thread_pool_exec(self.web_search.retrieve_chunks, query) except Exception: logging.exception("web_retrieve failed") return {"chunks": [], "doc_aggs": []} - return {"chunks": tav_res.get("chunks", []), "doc_aggs": tav_res.get("doc_aggs", [])} + return {"chunks": web_res.get("chunks", []), "doc_aggs": web_res.get("doc_aggs", [])} async def structured_retrieve(self, question: str) -> dict[str, Any]: """Query the structured (tabular) KBs by translating to SQL. diff --git a/rag/advanced_rag/harness/tools/search.py b/rag/advanced_rag/harness/tools/search.py index 492a6d6583..4e2d90b5bc 100644 --- a/rag/advanced_rag/harness/tools/search.py +++ b/rag/advanced_rag/harness/tools/search.py @@ -768,8 +768,8 @@ async def web_search(tools, query: str, keywords: str = "") -> dict: from common.misc_utils import thread_pool_exec effective_query = f"{query} {keywords}".strip() if keywords else query - tav_res = await thread_pool_exec(tools.tav.retrieve_chunks, effective_query) - return {"chunks": tav_res.get("chunks", []), "doc_aggs": tav_res.get("doc_aggs", [])} + web_res = await thread_pool_exec(tools.web_search.retrieve_chunks, effective_query) + return {"chunks": web_res.get("chunks", []), "doc_aggs": web_res.get("doc_aggs", [])} except Exception: _LOG.exception("web_search failed") return {"chunks": [], "doc_aggs": []} diff --git a/rag/advanced_rag/tree_structured_query_decomposition_retrieval.py b/rag/advanced_rag/tree_structured_query_decomposition_retrieval.py index c814bdee51..3b481b713a 100644 --- a/rag/advanced_rag/tree_structured_query_decomposition_retrieval.py +++ b/rag/advanced_rag/tree_structured_query_decomposition_retrieval.py @@ -19,7 +19,7 @@ from functools import partial from api.db.services.llm_service import LLMBundle from rag.prompts import kb_prompt from rag.prompts.generator import sufficiency_check, multi_queries_gen -from rag.utils.tavily_conn import Tavily +from rag.utils.web_search_conn import create_web_search_provider from timeit import default_timer as timer @@ -49,13 +49,13 @@ class TreeStructuredQueryDecompositionRetrieval: except Exception as e: logging.error(f"Knowledge base retrieval error: {e}") - # 2. Web retrieval (if Tavily API is configured) + # 2. Web retrieval (if a web search provider is configured) try: - if self.internet_enabled and self.prompt_config.get("tavily_api_key"): - tav = Tavily(self.prompt_config["tavily_api_key"]) - tav_res = tav.retrieve_chunks(search_query) - kbinfos["chunks"].extend(tav_res["chunks"]) - kbinfos["doc_aggs"].extend(tav_res["doc_aggs"]) + web_search = create_web_search_provider(self.prompt_config) if self.internet_enabled else None + if web_search: + web_res = web_search.retrieve_chunks(search_query) + kbinfos["chunks"].extend(web_res["chunks"]) + kbinfos["doc_aggs"].extend(web_res["doc_aggs"]) except Exception as e: logging.error(f"Web retrieval error: {e}") diff --git a/rag/utils/querit_conn.py b/rag/utils/querit_conn.py new file mode 100644 index 0000000000..030aae58b4 --- /dev/null +++ b/rag/utils/querit_conn.py @@ -0,0 +1,125 @@ +# +# Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +import logging +from typing import Any + +import requests + +from common.http_client import DEFAULT_TIMEOUT +from common.misc_utils import get_uuid +from rag.nlp import rag_tokenizer + +logger = logging.getLogger(__name__) + +QUERIT_SEARCH_URL = "https://api.querit.ai/v1/search" + + +class Querit: + def __init__(self, api_key: str): + self.api_key = api_key + + def search(self, query: str) -> list[dict[str, Any]]: + try: + response = requests.post( + QUERIT_SEARCH_URL, + headers={ + "Accept": "application/json", + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + }, + json={ + "query": query, + "count": 6, + "chunksPerDoc": 1, + }, + timeout=DEFAULT_TIMEOUT, + ) + response.raise_for_status() + response_data = response.json() + if not isinstance(response_data, dict): + raise TypeError("Querit API response must be a JSON object.") + + results_container = response_data.get("results", {}) + if not isinstance(results_container, dict): + raise TypeError("Querit API response field results must be an object.") + results = results_container.get("result", []) + if not isinstance(results, list): + raise TypeError("Querit API response field results.result must be an array.") + + normalized_results = [] + for result in results: + if not isinstance(result, dict): + continue + content = _querit_text(result.get("snippet")) + if not content: + continue + normalized_results.append( + { + "url": _querit_text(result.get("url")), + "title": _querit_text(result.get("title")), + "content": content, + "score": 1.0, + } + ) + return normalized_results + except (requests.RequestException, TypeError, ValueError) as error: + logger.error("Querit search failed: %s", _safe_error_message(error, self.api_key)) + return [] + + def retrieve_chunks(self, question: str) -> dict[str, list]: + chunks = [] + doc_aggs = [] + logger.info("[Querit]Q: %s", question) + for result in self.search(question): + chunk_id = get_uuid() + chunks.append( + { + "chunk_id": chunk_id, + "content_ltks": rag_tokenizer.tokenize(result["content"]), + "content_with_weight": result["content"], + "doc_id": chunk_id, + "docnm_kwd": result["title"], + "kb_id": [], + "important_kwd": [], + "image_id": "", + "similarity": result["score"], + "vector_similarity": 1.0, + "term_similarity": 0, + "vector": [], + "positions": [], + "url": result["url"], + } + ) + doc_aggs.append( + { + "doc_name": result["title"], + "doc_id": chunk_id, + "count": 1, + "url": result["url"], + } + ) + logger.info("[Querit]R: %s...", result["content"][:128]) + return {"chunks": chunks, "doc_aggs": doc_aggs} + + +def _querit_text(value: Any) -> str: + return "" if value is None else str(value) + + +def _safe_error_message(error: Exception, api_key: str) -> str: + message = str(error) or error.__class__.__name__ + return message.replace(api_key, "[REDACTED]") if api_key else message diff --git a/rag/utils/web_search_conn.py b/rag/utils/web_search_conn.py new file mode 100644 index 0000000000..2739c3eab3 --- /dev/null +++ b/rag/utils/web_search_conn.py @@ -0,0 +1,66 @@ +# +# Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +import logging +from typing import Protocol + +from rag.utils.querit_conn import Querit +from rag.utils.tavily_conn import Tavily + +WEB_SEARCH_PROVIDER_TAVILY = "tavily" +WEB_SEARCH_PROVIDER_QUERIT = "querit" + +logger = logging.getLogger(__name__) + + +class WebSearchProvider(Protocol): + def retrieve_chunks(self, question: str) -> dict[str, list]: + """Return web results in RAGFlow's chunk and document aggregate shape.""" + + +def _get_api_key(prompt_config: dict, field: str) -> str: + api_key = prompt_config.get(field) + return api_key.strip() if isinstance(api_key, str) else "" + + +def has_web_search_provider(prompt_config: dict | None) -> bool: + if not prompt_config: + return False + provider = prompt_config.get("web_search_provider", WEB_SEARCH_PROVIDER_TAVILY) + if provider == WEB_SEARCH_PROVIDER_TAVILY: + return bool(_get_api_key(prompt_config, "tavily_api_key")) + if provider == WEB_SEARCH_PROVIDER_QUERIT: + return bool(_get_api_key(prompt_config, "querit_api_key")) + return False + + +def create_web_search_provider(prompt_config: dict | None) -> WebSearchProvider | None: + if not prompt_config: + logger.debug("Web search provider resolution: provider=none status=disabled") + return None + + provider = prompt_config.get("web_search_provider", WEB_SEARCH_PROVIDER_TAVILY) + if provider not in (WEB_SEARCH_PROVIDER_TAVILY, WEB_SEARCH_PROVIDER_QUERIT): + logger.debug("Web search provider resolution: provider=%s status=invalid", provider) + return None + if not has_web_search_provider(prompt_config): + logger.debug("Web search provider resolution: provider=%s status=disabled", provider) + return None + + logger.debug("Web search provider resolution: provider=%s status=resolved", provider) + if provider == WEB_SEARCH_PROVIDER_QUERIT: + return Querit(_get_api_key(prompt_config, "querit_api_key")) + return Tavily(_get_api_key(prompt_config, "tavily_api_key")) diff --git a/test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py b/test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py index d811443d7e..379f0a31c8 100644 --- a/test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py +++ b/test/testcases/test_http_api/test_session_management/test_session_sdk_routes_unit.py @@ -1668,6 +1668,25 @@ def test_chatbot_routes_auth_stream_nonstream_unit(monkeypatch): assert res["data"]["avatar"] == "avatar.png" assert res["data"]["prologue"] == "Hello!" assert res["data"]["has_tavily_key"] is True + assert res["data"]["has_web_search_provider"] is True + + # Explicit Querit configuration also enables the provider-neutral flag. + querit_dialog = SimpleNamespace( + name="My Querit Bot", + icon="avatar.png", + tenant_id="tenant-1", + status="1", + llm_id="", + prompt_config={ + "prologue": "Hello!", + "web_search_provider": "querit", + "querit_api_key": "querit-key123", + }, + ) + monkeypatch.setattr(module.DialogService, "get_by_id", lambda _dialog_id: (True, querit_dialog)) + res = _run(inspect.unwrap(module.chatbots_inputs)("dialog-querit")) + assert res["code"] == 0 + assert res["data"]["has_web_search_provider"] is True @pytest.mark.p2 diff --git a/test/unit_test/rag/utils/test_querit_conn.py b/test/unit_test/rag/utils/test_querit_conn.py new file mode 100644 index 0000000000..40224a4778 --- /dev/null +++ b/test/unit_test/rag/utils/test_querit_conn.py @@ -0,0 +1,123 @@ +# +# Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +from rag.utils import querit_conn + + +class _Response: + status_code = 200 + + def raise_for_status(self): + return None + + def json(self): + return { + "results": { + "result": [ + { + "title": "RAGFlow", + "url": "https://example.com/ragflow", + "snippet": "RAGFlow is an open-source RAG engine.", + } + ] + } + } + + +def test_querit_search_uses_chat_defaults_and_normalizes_results(monkeypatch): + request = {} + + def fake_post(url, *, headers, json, timeout): + request.update(url=url, headers=headers, json=json, timeout=timeout) + return _Response() + + monkeypatch.setattr(querit_conn.requests, "post", fake_post) + + results = querit_conn.Querit("querit-test").search("What is RAGFlow?") + + assert request["url"] == "https://api.querit.ai/v1/search" + assert request["headers"]["Authorization"] == "Bearer querit-test" + assert request["json"] == { + "query": "What is RAGFlow?", + "count": 6, + "chunksPerDoc": 1, + } + assert results == [ + { + "url": "https://example.com/ragflow", + "title": "RAGFlow", + "content": "RAGFlow is an open-source RAG engine.", + "score": 1.0, + } + ] + + +def test_querit_retrieve_chunks_returns_ragflow_reference_shape(monkeypatch): + monkeypatch.setattr( + querit_conn.Querit, + "search", + lambda _self, _question: [ + { + "url": "https://example.com/ragflow", + "title": "RAGFlow", + "content": "RAGFlow is an open-source RAG engine.", + "score": 1.0, + } + ], + ) + monkeypatch.setattr(querit_conn, "get_uuid", lambda: "chunk-1") + monkeypatch.setattr(querit_conn.rag_tokenizer, "tokenize", lambda content: f"tokens:{content}") + + result = querit_conn.Querit("querit-test").retrieve_chunks("What is RAGFlow?") + + assert result["chunks"] == [ + { + "chunk_id": "chunk-1", + "content_ltks": "tokens:RAGFlow is an open-source RAG engine.", + "content_with_weight": "RAGFlow is an open-source RAG engine.", + "doc_id": "chunk-1", + "docnm_kwd": "RAGFlow", + "kb_id": [], + "important_kwd": [], + "image_id": "", + "similarity": 1.0, + "vector_similarity": 1.0, + "term_similarity": 0, + "vector": [], + "positions": [], + "url": "https://example.com/ragflow", + } + ] + assert result["doc_aggs"] == [ + { + "doc_name": "RAGFlow", + "doc_id": "chunk-1", + "count": 1, + "url": "https://example.com/ragflow", + } + ] + + +def test_querit_search_redacts_api_key_from_failures(monkeypatch, caplog): + class _FailedResponse: + def raise_for_status(self): + raise ValueError("request failed with querit-secret") + + monkeypatch.setattr(querit_conn.requests, "post", lambda *_args, **_kwargs: _FailedResponse()) + + assert querit_conn.Querit("querit-secret").search("RAGFlow") == [] + assert "querit-secret" not in caplog.text + assert "[REDACTED]" in caplog.text diff --git a/test/unit_test/rag/utils/test_web_search_conn.py b/test/unit_test/rag/utils/test_web_search_conn.py new file mode 100644 index 0000000000..98a09b4cdc --- /dev/null +++ b/test/unit_test/rag/utils/test_web_search_conn.py @@ -0,0 +1,101 @@ +# +# Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +from rag.utils import web_search_conn + + +def test_create_web_search_provider_uses_existing_tavily_config_without_provider_field(monkeypatch): + created_with = [] + provider = object() + + monkeypatch.setattr(web_search_conn, "Tavily", lambda api_key: created_with.append(api_key) or provider) + + result = web_search_conn.create_web_search_provider({"tavily_api_key": "tvly-test"}) + + assert result is provider + assert created_with == ["tvly-test"] + + +def test_create_web_search_provider_uses_selected_querit_config(monkeypatch): + created_with = [] + provider = object() + + monkeypatch.setattr(web_search_conn, "Querit", lambda api_key: created_with.append(api_key) or provider) + + result = web_search_conn.create_web_search_provider( + { + "web_search_provider": "querit", + "querit_api_key": "querit-test", + "tavily_api_key": "tvly-test", + } + ) + + assert result is provider + assert created_with == ["querit-test"] + + +def test_create_web_search_provider_trims_selected_key(monkeypatch): + created_with = [] + provider = object() + + monkeypatch.setattr(web_search_conn, "Querit", lambda api_key: created_with.append(api_key) or provider) + + result = web_search_conn.create_web_search_provider( + { + "web_search_provider": "querit", + "querit_api_key": " querit-test ", + } + ) + + assert result is provider + assert created_with == ["querit-test"] + + +def test_create_web_search_provider_requires_key_for_selected_provider(): + assert web_search_conn.create_web_search_provider({}) is None + assert web_search_conn.create_web_search_provider(None) is None + assert web_search_conn.create_web_search_provider({"web_search_provider": "tavily"}) is None + assert web_search_conn.create_web_search_provider({"web_search_provider": "querit"}) is None + assert web_search_conn.create_web_search_provider({"tavily_api_key": " "}) is None + assert ( + web_search_conn.create_web_search_provider( + { + "web_search_provider": "querit", + "querit_api_key": " ", + } + ) + is None + ) + + +def test_has_web_search_provider_follows_selected_provider(): + assert web_search_conn.has_web_search_provider({"tavily_api_key": "tvly-test"}) + assert not web_search_conn.has_web_search_provider({"tavily_api_key": ""}) + assert web_search_conn.has_web_search_provider({"web_search_provider": "querit", "querit_api_key": "querit-test"}) + assert not web_search_conn.has_web_search_provider( + { + "web_search_provider": "querit", + "querit_api_key": "", + "tavily_api_key": "tvly-test", + } + ) + assert not web_search_conn.has_web_search_provider( + { + "web_search_provider": "unsupported", + "querit_api_key": "querit-test", + "tavily_api_key": "tvly-test", + } + ) diff --git a/web/src/assets/querit.png b/web/src/assets/querit.png index 7a3045f625..223524516a 100644 Binary files a/web/src/assets/querit.png and b/web/src/assets/querit.png differ diff --git a/web/src/components/tavily-form-field.tsx b/web/src/components/tavily-form-field.tsx deleted file mode 100644 index abccae52bb..0000000000 --- a/web/src/components/tavily-form-field.tsx +++ /dev/null @@ -1,51 +0,0 @@ -import { useTranslate } from '@/hooks/common-hooks'; -import { useFormContext } from 'react-hook-form'; -import PasswordInput from './originui/password-input'; -import { - FormControl, - FormDescription, - FormField, - FormItem, - FormLabel, - FormMessage, -} from './ui/form'; - -interface IProps { - name?: string; -} - -export function TavilyFormField({ - name = 'prompt_config.tavily_api_key', -}: IProps) { - const form = useFormContext(); - const { t } = useTranslate('chat'); - - return ( - ( - - Tavily API Key - - - - - - {t('tavilyApiKeyHelp')} - - - - - )} - /> - ); -} diff --git a/web/src/components/web-search-form-field.tsx b/web/src/components/web-search-form-field.tsx new file mode 100644 index 0000000000..eccace2cf4 --- /dev/null +++ b/web/src/components/web-search-form-field.tsx @@ -0,0 +1,131 @@ +import queritLogo from '@/assets/querit.png'; +import tavilyLogo from '@/assets/svg/tavily.svg'; +import { RAGFlowSelect } from '@/components/ui/select'; +import { WebSearchProvider } from '@/constants/chat'; +import { useTranslate } from '@/hooks/common-hooks'; +import { prefixName } from '@/utils/form'; +import { useFormContext, useWatch } from 'react-hook-form'; +import PasswordInput from './originui/password-input'; +import { + FormControl, + FormDescription, + FormField, + FormItem, + FormLabel, + FormMessage, +} from './ui/form'; + +interface IProps { + prefix?: string; +} + +const providerOptions = [ + { + name: 'Tavily', + logo: tavilyLogo, + value: WebSearchProvider.Tavily, + }, + { + name: 'Querit', + logo: queritLogo, + value: WebSearchProvider.Querit, + }, +] + .sort((left, right) => left.name.localeCompare(right.name)) + .map(({ name, logo, value }) => ({ + label: ( + + + {name} + + ), + value, + })); + +const providerKeyConfig = { + [WebSearchProvider.Tavily]: { + name: 'prompt_config.tavily_api_key', + label: 'Tavily API Key', + tip: 'tavilyApiKeyTip', + placeholder: 'tavilyApiKeyMessage', + helpUrl: 'https://app.tavily.com/home', + }, + [WebSearchProvider.Querit]: { + name: 'prompt_config.querit_api_key', + label: 'Querit API Key', + tip: 'queritApiKeyTip', + placeholder: 'queritApiKeyMessage', + helpUrl: 'https://querit.ai', + }, +} as const; + +export function WebSearchFormField({ prefix = '' }: IProps) { + const form = useFormContext(); + const { t } = useTranslate('chat'); + const providerName = prefixName(prefix, 'prompt_config.web_search_provider'); + const selectedProvider = useWatch({ + control: form.control, + name: providerName, + }); + const keyConfig = providerKeyConfig[selectedProvider as WebSearchProvider]; + + return ( + <> + ( + + + {t('webSearchProvider')} + + + + + + + )} + /> + {keyConfig && ( + ( + + + {keyConfig.label} + + + + + + + {t('tavilyApiKeyHelp')} + + + + + )} + /> + )} + + ); +} diff --git a/web/src/constants/chat.ts b/web/src/constants/chat.ts index f102d9e79b..0570474140 100644 --- a/web/src/constants/chat.ts +++ b/web/src/constants/chat.ts @@ -39,3 +39,8 @@ export enum DatasetMetadata { SemiAutomatic = 'semi_auto', Manual = 'manual', } + +export enum WebSearchProvider { + Tavily = 'tavily', + Querit = 'querit', +} diff --git a/web/src/interfaces/database/chat.ts b/web/src/interfaces/database/chat.ts index 2bd8d4c3a0..f5dc1f7318 100644 --- a/web/src/interfaces/database/chat.ts +++ b/web/src/interfaces/database/chat.ts @@ -1,4 +1,4 @@ -import { MessageType } from '@/constants/chat'; +import { MessageType, WebSearchProvider } from '@/constants/chat'; import { IAttachment } from '@/hooks/use-send-message'; export interface IDocumentDownloadInfo { @@ -21,6 +21,8 @@ export interface PromptConfig { reasoning?: boolean; cross_languages?: Array; tavily_api_key?: string; + querit_api_key?: string; + web_search_provider?: WebSearchProvider; toc_enhance?: boolean; reference_metadata?: { include?: boolean; @@ -202,6 +204,7 @@ export interface IExternalChatInfo { title: string; prologue?: string; has_tavily_key?: boolean; + has_web_search_provider?: boolean; llm_id?: string; } diff --git a/web/src/locales/en.ts b/web/src/locales/en.ts index b157c4a063..e72b6e0e5a 100644 --- a/web/src/locales/en.ts +++ b/web/src/locales/en.ts @@ -1244,6 +1244,13 @@ This auto-tagging feature enhances retrieval by adding another layer of domain-s tavilyApiKeyTip: 'If an API Key is correctly set here, Tavily-based web searches will be used to supplement dataset retrieval.', tavilyApiKeyMessage: 'Please enter your Tavily API Key', + webSearchProvider: 'Web search provider', + webSearchProviderTip: + 'Select the service used when Internet search is enabled.', + webSearchProviderPlaceholder: 'Select a web search provider', + queritApiKeyTip: + 'When Querit is selected, its web search results supplement dataset retrieval.', + queritApiKeyMessage: 'Please enter your Querit API Key', tavilyApiKeyHelp: 'How to get it?', crossLanguage: 'Cross-language search', crossLanguagePlaceholder: 'Select value', diff --git a/web/src/locales/zh.ts b/web/src/locales/zh.ts index 4f1d5995a7..1076886997 100644 --- a/web/src/locales/zh.ts +++ b/web/src/locales/zh.ts @@ -1131,6 +1131,12 @@ NER:使用 spaCy NER 和基于规则的关键词提取来抽取实体和关系 tavilyApiKeyTip: '如果 API 密钥设置正确,它将利用 Tavily 进行网络搜索作为知识库的补充。', tavilyApiKeyMessage: '请输入你的 Tavily API Key', + webSearchProvider: '网络搜索服务', + webSearchProviderTip: '选择启用联网搜索时使用的搜索服务。', + webSearchProviderPlaceholder: '请选择网络搜索服务', + queritApiKeyTip: + '选择 Querit 后,将使用 Querit 的网络搜索结果补充知识库检索。', + queritApiKeyMessage: '请输入你的 Querit API Key', tavilyApiKeyHelp: '如何获取?', crossLanguage: '跨语言搜索', crossLanguagePlaceholder: '请选择', diff --git a/web/src/pages/next-chats/chat/app-settings/chat-prompt-engine.tsx b/web/src/pages/next-chats/chat/app-settings/chat-prompt-engine.tsx index e7eb59642c..921dccec0c 100644 --- a/web/src/pages/next-chats/chat/app-settings/chat-prompt-engine.tsx +++ b/web/src/pages/next-chats/chat/app-settings/chat-prompt-engine.tsx @@ -6,8 +6,6 @@ import { MetadataFilter } from '@/components/metadata-filter'; import { RerankFormFields } from '@/components/rerank'; import { SimilaritySliderFormField } from '@/components/similarity-slider'; import { SwitchFormField } from '@/components/switch-fom-field'; -import { TavilyFormField } from '@/components/tavily-form-field'; - import { TopNFormField } from '@/components/top-n-item'; import { FormControl, @@ -19,7 +17,7 @@ import { import { MultiSelect } from '@/components/ui/multi-select'; import { Switch } from '@/components/ui/switch'; import { Textarea } from '@/components/ui/textarea'; - +import { WebSearchFormField } from '@/components/web-search-form-field'; import { useFetchKnowledgeMetadataKeys } from '@/hooks/use-knowledge-request'; import { prefixName } from '@/utils/form'; import { getDirAttribute } from '@/utils/text-direction'; @@ -127,9 +125,7 @@ export function ChatPromptEngine({ prefix = '' }: ChatPromptEngineProps) { label={t('chat.tts')} tooltip={t('chat.ttsTip')} > - + { + it('does not select a provider for a new unconfigured dialog', () => { + expect(getWebSearchProvider({} as PromptConfig)).toBeUndefined(); + }); + + it('selects Tavily for a legacy dialog with a Tavily key', () => { + const promptConfig = { + tavily_api_key: 'tvly-test', + } as PromptConfig; + + expect(getWebSearchProvider(promptConfig)).toBe(WebSearchProvider.Tavily); + }); +}); + +describe('getWebSearchApiKey', () => { + it('uses Tavily for dialogs saved before provider selection existed', () => { + const promptConfig = { + tavily_api_key: 'tvly-test', + } as PromptConfig; + + expect(getWebSearchApiKey(promptConfig)).toBe('tvly-test'); + }); + + it('uses only the selected Querit key', () => { + const promptConfig = { + web_search_provider: WebSearchProvider.Querit, + querit_api_key: 'querit-test', + tavily_api_key: 'tvly-test', + } as PromptConfig; + + expect(getWebSearchApiKey(promptConfig)).toBe('querit-test'); + }); + + it('does not fall back to Tavily when Querit is selected without a key', () => { + const promptConfig = { + web_search_provider: WebSearchProvider.Querit, + tavily_api_key: 'tvly-test', + } as PromptConfig; + + expect(getWebSearchApiKey(promptConfig)).toBeUndefined(); + }); + + it('treats a whitespace-only key as unconfigured', () => { + const promptConfig = { + web_search_provider: WebSearchProvider.Querit, + querit_api_key: ' ', + } as PromptConfig; + + expect(getWebSearchApiKey(promptConfig)).toBe(''); + }); + + it('does not fall back to Tavily for an unsupported provider', () => { + const promptConfig = { + web_search_provider: 'unsupported', + tavily_api_key: 'tvly-test', + } as unknown as PromptConfig; + + expect(getWebSearchApiKey(promptConfig)).toBeUndefined(); + }); + + it('treats a non-string key as unconfigured', () => { + const promptConfig = { + web_search_provider: WebSearchProvider.Querit, + querit_api_key: 123, + } as unknown as PromptConfig; + + expect(getWebSearchApiKey(promptConfig)).toBeUndefined(); + }); +}); diff --git a/web/src/pages/next-chats/chat/use-show-internet.ts b/web/src/pages/next-chats/chat/use-show-internet.ts index 64ac48c20e..3db50dd0a9 100644 --- a/web/src/pages/next-chats/chat/use-show-internet.ts +++ b/web/src/pages/next-chats/chat/use-show-internet.ts @@ -1,8 +1,9 @@ import { useFetchChat } from '@/hooks/use-chat-request'; import { isEmpty } from 'lodash'; +import { getWebSearchApiKey } from './web-search-api-key'; export function useShowInternet() { const { data: currentDialog } = useFetchChat(); - return !isEmpty(currentDialog?.prompt_config?.tavily_api_key); + return !isEmpty(getWebSearchApiKey(currentDialog?.prompt_config)); } diff --git a/web/src/pages/next-chats/chat/web-search-api-key.ts b/web/src/pages/next-chats/chat/web-search-api-key.ts new file mode 100644 index 0000000000..554b94d5c4 --- /dev/null +++ b/web/src/pages/next-chats/chat/web-search-api-key.ts @@ -0,0 +1,41 @@ +import { WebSearchProvider } from '@/constants/chat'; +import type { PromptConfig } from '@/interfaces/database/chat'; + +export function getWebSearchProvider(promptConfig?: PromptConfig) { + const provider = promptConfig?.web_search_provider; + + if ( + provider === WebSearchProvider.Tavily || + provider === WebSearchProvider.Querit + ) { + return provider; + } + + if ( + provider === undefined && + typeof promptConfig?.tavily_api_key === 'string' && + promptConfig.tavily_api_key.trim() + ) { + return WebSearchProvider.Tavily; + } + + return undefined; +} + +export function getWebSearchApiKey(promptConfig?: PromptConfig) { + const provider = getWebSearchProvider(promptConfig); + let apiKey: unknown; + + switch (provider) { + case WebSearchProvider.Tavily: + apiKey = promptConfig?.tavily_api_key; + break; + case WebSearchProvider.Querit: + apiKey = promptConfig?.querit_api_key; + break; + default: + return undefined; + } + + return typeof apiKey === 'string' ? apiKey.trim() : undefined; +} diff --git a/web/src/pages/next-chats/share/index.tsx b/web/src/pages/next-chats/share/index.tsx index 96c44ea463..8a431acc1c 100644 --- a/web/src/pages/next-chats/share/index.tsx +++ b/web/src/pages/next-chats/share/index.tsx @@ -115,7 +115,9 @@ const ChatContainer = () => { showUploadIcon={false} stopOutputMessage={stopOutputMessage} showReasoning - showInternet={chatInfo?.has_tavily_key} + showInternet={ + chatInfo?.has_web_search_provider ?? chatInfo?.has_tavily_key + } >