feat(agent): add Querit search tool (#17548)

This commit is contained in:
EthanZhang
2026-07-30 09:36:16 +08:00
committed by GitHub
parent 03d887a975
commit 33e581a8b3
45 changed files with 3527 additions and 26 deletions

View File

@@ -20,6 +20,7 @@ import (
"context"
"encoding/json"
"fmt"
"io"
"strings"
"ragflow/internal/agent/runtime"
@@ -69,7 +70,12 @@ func (c *ToolBackedComponent) Invoke(ctx context.Context, db *gorm.DB, inputs ma
}
raw, invokeErr := c.tool.InvokableRun(ctx, string(argsJSON))
decoded := parseToolEnvelope(raw)
var decoded map[string]any
if c.spec.PreserveJSONNumbers {
decoded = parseToolEnvelopeLossless(raw)
} else {
decoded = parseToolEnvelope(raw)
}
if rawValue, invalid := decoded["_raw"]; invalid {
if invokeErr != nil {
return nil, fmt.Errorf("canvas: %s: %w", c.name, invokeErr)
@@ -97,6 +103,19 @@ func (c *ToolBackedComponent) Invoke(ctx context.Context, db *gorm.DB, inputs ma
return c.tool.BuildComponentOutputs(decoded), nil
}
func parseToolEnvelopeLossless(jsonStr string) map[string]any {
var out map[string]any
decoder := json.NewDecoder(strings.NewReader(jsonStr))
decoder.UseNumber()
if err := decoder.Decode(&out); err != nil {
return map[string]any{"_raw": jsonStr}
}
if err := decoder.Decode(&struct{}{}); err != io.EOF {
return map[string]any{"_raw": jsonStr}
}
return out
}
func (c *ToolBackedComponent) Stream(_ context.Context, _ *gorm.DB, _ map[string]any) (<-chan map[string]any, error) {
return nil, nil
}
@@ -115,6 +134,7 @@ var toolComponentRegistrations = []struct {
{componentName: "GoogleScholar", toolName: "google_scholar"},
{componentName: "KeenableSearch", toolName: "keenable"},
{componentName: "PubMed", toolName: "pubmed"},
{componentName: "QueritSearch", toolName: "querit_search"},
{componentName: "SearXNG", toolName: "searxng"},
{componentName: "TavilySearch", toolName: "tavily"},
{componentName: "TavilyExtract", toolName: "tavily_extract"},

View File

@@ -25,6 +25,7 @@ import (
"net/url"
"reflect"
"strings"
"sync/atomic"
"testing"
einotool "github.com/cloudwego/eino/components/tool"
@@ -192,6 +193,13 @@ func TestToolBackedComponentRegisteredFactories(t *testing.T) {
outputKey: "success",
inputKey: "to_email",
},
{
name: "QueritSearch",
toolName: "QueritSearch",
params: map[string]any{"api_key": "stored-key", "count": float64(10), "chunks_per_doc": float64(3), "outputs": map[string]any{"json": map[string]any{}}},
outputKey: "json",
inputKey: "query",
},
{
name: "SearXNG",
toolName: "SearXNG",
@@ -275,7 +283,7 @@ func TestToolBackedComponentWenCaiInvoke(t *testing.T) {
}
func TestToolBackedComponentRegisteredBuildWorkflow(t *testing.T) {
for _, componentName := range []string{"ArXiv", "BGPT", "DuckDuckGo", "Email", "Google", "GoogleScholar", "KeenableSearch", "PubMed", "SearXNG", "WenCai", "TavilyExtract", "TavilySearch", "Wikipedia", "YahooFinance"} {
for _, componentName := range []string{"ArXiv", "BGPT", "DuckDuckGo", "Email", "Google", "GoogleScholar", "KeenableSearch", "PubMed", "QueritSearch", "SearXNG", "WenCai", "TavilyExtract", "TavilySearch", "Wikipedia", "YahooFinance"} {
t.Run(componentName, func(t *testing.T) {
c := &canvas.Canvas{
Components: map[string]canvas.CanvasComponent{
@@ -468,6 +476,65 @@ func TestToolBackedComponentTavilyIntegration(t *testing.T) {
}
}
func TestToolBackedComponentQueritIntegration(t *testing.T) {
var serverCalls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
serverCalls.Add(1)
writer.Header().Set("Content-Type", "application/json")
_, _ = writer.Write([]byte(`{"results":{"result":[{"title":"RAGFlow","url":"https://ragflow.io","snippet":"RAG article","custom":"preserved"}]},"search_id":11099848653006015581,"query_context":{"rewritten":"rag flow"}}`))
}))
defer server.Close()
target, err := url.Parse(server.URL)
if err != nil {
t.Fatalf("parse test server URL: %v", err)
}
helper := agenttool.NewHTTPHelper().WithClient(&http.Client{Transport: roundTripperFunc(func(request *http.Request) (*http.Response, error) {
cloned := request.Clone(request.Context())
cloned.URL.Scheme = target.Scheme
cloned.URL.Host = target.Host
return http.DefaultTransport.RoundTrip(cloned)
})})
querit := agenttool.NewQueritToolWithEnvKey(helper, func() string { return "key" })
component := &ToolBackedComponent{name: "QueritSearch", tool: querit, spec: querit.ComponentSpec()}
state := runtime.NewCanvasState("run-querit", "task-querit")
empty, err := component.Invoke(context.Background(), nil, map[string]any{"query": ""})
if err != nil {
t.Fatalf("Invoke(empty query): %v", err)
}
emptyJSON, emptyJSONOK := empty["json"].(map[string]any)
if serverCalls.Load() != 0 || empty["formalized_content"] != "" || !emptyJSONOK || len(emptyJSON) != 0 {
t.Fatalf("empty query result = %#v, server calls = %d", empty, serverCalls.Load())
}
out, err := component.Invoke(runtime.WithState(context.Background(), state), nil, map[string]any{"query": "ragflow"})
if err != nil {
t.Fatalf("Invoke: %v", err)
}
rendered, _ := out["formalized_content"].(string)
if !strings.Contains(rendered, "Title: RAGFlow") || !strings.Contains(rendered, "RAG article") {
t.Fatalf("formalized_content = %q", rendered)
}
complete, ok := out["json"].(map[string]any)
searchID, searchIDOK := complete["search_id"].(json.Number)
if !ok || !searchIDOK || searchID.String() != "11099848653006015581" || complete["query_context"] == nil {
t.Fatalf("complete json = %#v", out["json"])
}
encoded, err := json.Marshal(complete)
if err != nil || !strings.Contains(string(encoded), `"search_id":11099848653006015581`) {
t.Fatalf("re-encoded complete json = %s, %v", encoded, err)
}
chunks := state.GetRetrievalChunks()
if len(chunks) != 1 || chunks[0]["document_name"] != "RAGFlow" || chunks[0]["similarity"] != 1 {
t.Fatalf("recorded references = %#v", chunks)
}
}
func TestParseToolEnvelopeLosslessRejectsTrailingContent(t *testing.T) {
out := parseToolEnvelopeLossless(`{"results":{"result":[]}} trailing`)
if out["_raw"] == nil {
t.Fatalf("trailing content was accepted: %#v", out)
}
}
func TestToolBackedComponentYahooFinanceIntegration(t *testing.T) {
serverCalls := 0
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {

View File

@@ -0,0 +1,605 @@
//
// 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 tool
import (
"context"
"crypto/sha1"
"encoding/json"
"fmt"
"io"
"math/big"
"net/http"
"regexp"
"strconv"
"strings"
"time"
"github.com/cloudwego/eino/components/tool"
"github.com/cloudwego/eino/schema"
"ragflow/internal/common"
"ragflow/internal/tokenizer"
)
const (
queritToolName = "querit_search"
queritToolDescription = "Search the web via the Querit API and return the complete search response."
queritEndpoint = "https://api.querit.ai/v1/search"
queritMaxAttempts = 3
queritPromptMaxTokens = 200000
)
var (
queritRelativeTimeRangePattern = regexp.MustCompile(`^[dwmy][1-9][0-9]*$`)
queritDateRangePattern = regexp.MustCompile(`^\d{4}-\d{2}-\d{2}to\d{4}-\d{2}-\d{2}$`)
queritDataImagePattern = regexp.MustCompile(`!?\[[a-z]+\]\(data:image/png;base64,[ 0-9A-Za-z/_=+\-]+\)`)
queritNewlinePattern = regexp.MustCompile(`\n+`)
)
type queritParams struct {
APIKey string `json:"api_key"`
Query string `json:"query"`
Count int `json:"count"`
ChunksPerDoc *int `json:"chunks_per_doc"`
SiteInclude []string `json:"site_include"`
SiteExclude []string `json:"site_exclude"`
TimeRange string `json:"time_range"`
CountryInclude []string `json:"country_include"`
LanguageInclude []string `json:"language_include"`
}
type queritRequest struct {
Query string `json:"query"`
Count int `json:"count"`
ChunksPerDoc *int `json:"chunksPerDoc,omitempty"`
Filters *queritFilters `json:"filters,omitempty"`
}
type queritFilters struct {
Sites *queritIncludeExcludeFilter `json:"sites,omitempty"`
TimeRange *queritDateFilter `json:"timeRange,omitempty"`
Geo *queritGeoFilter `json:"geo,omitempty"`
Languages *queritIncludeFilter `json:"languages,omitempty"`
}
type queritIncludeExcludeFilter struct {
Include []string `json:"include,omitempty"`
Exclude []string `json:"exclude,omitempty"`
}
type queritIncludeFilter struct {
Include []string `json:"include,omitempty"`
}
type queritDateFilter struct {
Date string `json:"date"`
}
type queritGeoFilter struct {
Countries *queritIncludeFilter `json:"countries,omitempty"`
}
// QueritTool searches the web through Querit. Node-level parameters are
// retained as defaults and model-emitted parameters override them per call.
type QueritTool struct {
helper *HTTPHelper
envKey func() string
defaults queritParams
retryWait func(context.Context, int) bool
}
var _ ToolComponent = (*QueritTool)(nil)
var _ ReferenceBuilder = (*QueritTool)(nil)
// NewQueritTool returns a QueritTool using the shared HTTP helper and the
// QUERIT_API_KEY environment variable.
func NewQueritTool() *QueritTool {
return newQueritTool(nil, nil, queritParams{}, nil)
}
// NewQueritToolWith returns a QueritTool using the supplied HTTP helper.
func NewQueritToolWith(helper *HTTPHelper) *QueritTool {
return newQueritTool(helper, nil, queritParams{}, nil)
}
// NewQueritToolWithEnvKey returns a QueritTool with an injectable API-key
// resolver. It is intended for tests that must not depend on process state.
func NewQueritToolWithEnvKey(helper *HTTPHelper, envKey func() string) *QueritTool {
return newQueritTool(helper, envKey, queritParams{}, nil)
}
func newQueritTool(
helper *HTTPHelper,
envKey func() string,
defaults queritParams,
retryWait func(context.Context, int) bool,
) *QueritTool {
if helper == nil {
helper = NewHTTPHelper()
}
if envKey == nil {
envKey = defaultQueritEnvKey
}
if defaults.Count == 0 {
defaults.Count = 10
}
if defaults.ChunksPerDoc == nil {
defaults.ChunksPerDoc = queritInt(3)
}
if retryWait == nil {
retryWait = waitForQueritRetry
}
return &QueritTool{helper: helper, envKey: envKey, defaults: defaults, retryWait: retryWait}
}
func defaultQueritEnvKey() string { return common.GetEnv(common.EnvQueritAPIKey) }
// Info returns the arguments that a chat model may provide. Credentials stay
// in node configuration or the environment and are never exposed to the model.
func (q *QueritTool) Info(_ context.Context) (*schema.ToolInfo, error) {
return &schema.ToolInfo{
Name: queritToolName,
Desc: queritToolDescription,
ParamsOneOf: schema.NewParamsOneOfByParams(map[string]*schema.ParameterInfo{
"query": {
Type: schema.String,
Desc: "Search query.",
Required: true,
},
"count": {
Type: schema.Integer,
Desc: "Maximum number of documents to return. Defaults to 10 and must be at least 1.",
Required: false,
},
"chunks_per_doc": {
Type: schema.Integer,
Desc: "Number of snippets per document. Defaults to 3 and must be between 1 and 3.",
Required: false,
},
"site_include": {
Type: schema.Array,
ElemInfo: &schema.ParameterInfo{Type: schema.String},
Desc: "Sites that search results must include.",
Required: false,
},
"site_exclude": {
Type: schema.Array,
ElemInfo: &schema.ParameterInfo{Type: schema.String},
Desc: "Sites that search results must exclude.",
Required: false,
},
"time_range": {
Type: schema.String,
Desc: "Relative time range such as d7, w1, m3, or y1, or YYYY-MM-DDtoYYYY-MM-DD.",
Required: false,
},
"country_include": {
Type: schema.Array,
ElemInfo: &schema.ParameterInfo{Type: schema.String},
Desc: "Countries that search results must include.",
Required: false,
},
"language_include": {
Type: schema.Array,
ElemInfo: &schema.ParameterInfo{Type: schema.String},
Desc: "Languages that search results must include.",
Required: false,
},
}),
}, nil
}
// InvokableRun performs a Querit search. All expected operational failures are
// returned as soft-error JSON with a nil Go error so an Agent run can continue.
func (q *QueritTool) InvokableRun(ctx context.Context, argsJSON string, _ ...tool.Option) (string, error) {
provided := make(map[string]json.RawMessage)
if err := json.Unmarshal([]byte(argsJSON), &provided); err != nil {
return queritErrorJSON(fmt.Errorf("parse arguments: %w", err)), nil
}
queryJSON, hasQuery := provided["query"]
if hasQuery {
if strings.TrimSpace(string(queryJSON)) == "null" {
return queritErrorJSON(fmt.Errorf("query must be provided as a string")), nil
}
var query string
if err := json.Unmarshal(queryJSON, &query); err != nil {
return queritErrorJSON(fmt.Errorf("query must be provided as a string")), nil
}
}
var runtimeParams queritParams
if err := json.Unmarshal([]byte(argsJSON), &runtimeParams); err != nil {
return queritErrorJSON(fmt.Errorf("parse arguments: %w", err)), nil
}
params := mergeQueritParams(q.defaults, runtimeParams, provided)
if !hasQuery && strings.TrimSpace(params.Query) == "" {
return queritErrorJSON(fmt.Errorf("query must be provided as a string")), nil
}
if params.Query == "" {
return `{}`, nil
}
if err := validateQueritParams(params); err != nil {
return queritErrorJSON(err, params.APIKey), nil
}
apiKey := strings.TrimSpace(params.APIKey)
if apiKey == "" {
apiKey = strings.TrimSpace(q.envKey())
}
if apiKey == "" {
return queritErrorJSON(fmt.Errorf("api_key is required (or set QUERIT_API_KEY)")), nil
}
body, err := json.Marshal(buildQueritRequest(params))
if err != nil {
return queritErrorJSON(fmt.Errorf("encode request: %w", err), apiKey), nil
}
for attempt := 1; attempt <= queritMaxAttempts; attempt++ {
resp, requestErr := q.helper.Do(
ctx,
http.MethodPost,
queritEndpoint,
string(body),
"application/json",
map[string]string{"Authorization": "Bearer " + apiKey},
)
if requestErr != nil {
return queritErrorJSON(fmt.Errorf("request failed: %w", requestErr), apiKey), nil
}
if resp.StatusCode == http.StatusTooManyRequests {
_, _ = io.Copy(io.Discard, resp.Body)
_ = resp.Body.Close()
if attempt == queritMaxAttempts || !q.retryWait(ctx, attempt) {
return queritErrorJSON(fmt.Errorf("upstream returned %d after %d attempts", resp.StatusCode, attempt), apiKey), nil
}
continue
}
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
_, _ = io.Copy(io.Discard, resp.Body)
_ = resp.Body.Close()
return queritErrorJSON(fmt.Errorf("upstream returned %d", resp.StatusCode), apiKey), nil
}
raw, readErr := io.ReadAll(resp.Body)
_ = resp.Body.Close()
if readErr != nil {
return queritErrorJSON(fmt.Errorf("read response: %w", readErr), apiKey), nil
}
if _, err := decodeQueritResponse(raw); err != nil {
return queritErrorJSON(err, apiKey), nil
}
return string(raw), nil
}
return queritErrorJSON(fmt.Errorf("request exhausted retries"), apiKey), nil
}
func mergeQueritParams(defaults, runtimeParams queritParams, provided map[string]json.RawMessage) queritParams {
merged := defaults
if _, ok := provided["query"]; ok {
merged.Query = runtimeParams.Query
}
if _, ok := provided["count"]; ok {
merged.Count = runtimeParams.Count
}
if _, ok := provided["chunks_per_doc"]; ok {
merged.ChunksPerDoc = runtimeParams.ChunksPerDoc
}
if _, ok := provided["site_include"]; ok {
merged.SiteInclude = runtimeParams.SiteInclude
}
if _, ok := provided["site_exclude"]; ok {
merged.SiteExclude = runtimeParams.SiteExclude
}
if _, ok := provided["time_range"]; ok {
merged.TimeRange = runtimeParams.TimeRange
}
if _, ok := provided["country_include"]; ok {
merged.CountryInclude = runtimeParams.CountryInclude
}
if _, ok := provided["language_include"]; ok {
merged.LanguageInclude = runtimeParams.LanguageInclude
}
return merged
}
func validateQueritParams(params queritParams) error {
if params.Count < 1 {
return fmt.Errorf("count must be at least 1")
}
if params.ChunksPerDoc != nil && (*params.ChunksPerDoc < 1 || *params.ChunksPerDoc > 3) {
return fmt.Errorf("chunks_per_doc must be between 1 and 3")
}
if !isValidQueritTimeRange(strings.TrimSpace(params.TimeRange)) {
return fmt.Errorf("time_range must use dN, wN, mN, yN, or YYYY-MM-DDtoYYYY-MM-DD")
}
return nil
}
func isValidQueritTimeRange(value string) bool {
return value == "" || queritRelativeTimeRangePattern.MatchString(value) || queritDateRangePattern.MatchString(value)
}
func buildQueritRequest(params queritParams) queritRequest {
request := queritRequest{
Query: params.Query,
Count: params.Count,
ChunksPerDoc: params.ChunksPerDoc,
}
filters := &queritFilters{}
if len(params.SiteInclude) > 0 || len(params.SiteExclude) > 0 {
filters.Sites = &queritIncludeExcludeFilter{Include: params.SiteInclude, Exclude: params.SiteExclude}
}
if strings.TrimSpace(params.TimeRange) != "" {
filters.TimeRange = &queritDateFilter{Date: params.TimeRange}
}
if len(params.CountryInclude) > 0 {
filters.Geo = &queritGeoFilter{Countries: &queritIncludeFilter{Include: params.CountryInclude}}
}
if len(params.LanguageInclude) > 0 {
filters.Languages = &queritIncludeFilter{Include: params.LanguageInclude}
}
if filters.Sites != nil || filters.TimeRange != nil || filters.Geo != nil || filters.Languages != nil {
request.Filters = filters
}
return request
}
func decodeQueritResponse(raw []byte) (map[string]any, error) {
var decoded any
decoder := json.NewDecoder(strings.NewReader(string(raw)))
decoder.UseNumber()
if err := decoder.Decode(&decoded); err != nil {
return nil, fmt.Errorf("decode response: %w", err)
}
if err := decoder.Decode(&struct{}{}); err != io.EOF {
if err == nil {
return nil, fmt.Errorf("decode response: expected exactly one JSON value")
}
return nil, fmt.Errorf("decode response: trailing content: %w", err)
}
response, ok := decoded.(map[string]any)
if !ok || response == nil {
return nil, fmt.Errorf("decode response: expected a JSON object")
}
resultsValue, exists := response["results"]
if !exists {
return response, nil
}
results, ok := resultsValue.(map[string]any)
if !ok {
return nil, fmt.Errorf("decode response: results must be a JSON object")
}
resultValue, exists := results["result"]
if !exists {
return response, nil
}
if _, ok := resultValue.([]any); !ok {
return nil, fmt.Errorf("decode response: results.result must be a JSON array")
}
return response, nil
}
func waitForQueritRetry(ctx context.Context, attempt int) bool {
delay := 200 * time.Millisecond
for current := 1; current < attempt; current++ {
delay *= 2
}
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-ctx.Done():
return false
case <-timer.C:
return true
}
}
// ComponentSpec returns the QueritSearch Canvas-facing metadata.
func (q *QueritTool) ComponentSpec() ComponentSpec {
return ComponentSpec{
PreserveJSONNumbers: true,
Inputs: map[string]string{
"api_key": "Querit API key. Uses QUERIT_API_KEY when empty.",
"query": "Search query.",
"count": "Maximum number of documents.",
"chunks_per_doc": "Number of snippets per document.",
"site_include": "Sites that search results must include.",
"site_exclude": "Sites that search results must exclude.",
"time_range": "dN, wN, mN, yN, or YYYY-MM-DDtoYYYY-MM-DD.",
"country_include": "Countries that search results must include.",
"language_include": "Languages that search results must include.",
},
Outputs: map[string]string{
"formalized_content": "Rendered Querit references for downstream prompts.",
"json": "Complete raw Querit JSON response.",
},
InputForm: map[string]any{
"api_key": map[string]any{"name": "API key", "type": "password"},
"query": map[string]any{"name": "Query", "type": "line"},
"count": map[string]any{"name": "Count", "type": "line"},
"chunks_per_doc": map[string]any{"name": "Chunks per document", "type": "line"},
"site_include": map[string]any{"name": "Site include", "type": "line"},
"site_exclude": map[string]any{"name": "Site exclude", "type": "line"},
"time_range": map[string]any{"name": "Time range", "type": "line"},
"country_include": map[string]any{"name": "Country include", "type": "line"},
"language_include": map[string]any{"name": "Language include", "type": "line"},
},
}
}
// BuildReferences converts results.result into RAGFlow retrieval references.
func (q *QueritTool) BuildReferences(_ context.Context, response map[string]any) ([]map[string]any, []map[string]any) {
results := queritResults(response)
chunks := make([]map[string]any, 0, len(results))
docAggs := make([]map[string]any, 0, len(results))
for _, result := range results {
title := queritText(result["title"])
resultURL := queritText(result["url"])
content := queritText(result["snippet"])
if content == "" {
continue
}
content = queritDataImagePattern.ReplaceAllString(content, "")
content = truncateQueritRunes(content, 10000)
if content == "" {
continue
}
documentID := strconv.FormatInt(queritHashInt(content, 100000000), 10)
displayID := strconv.FormatInt(queritHashInt(documentID, 500), 10)
chunks = append(chunks, map[string]any{
"id": displayID,
"chunk_id": documentID,
"content": content,
"doc_id": documentID,
"document_id": documentID,
"docnm_kwd": title,
"document_name": title,
"similarity": 1,
"score": 1,
"url": resultURL,
})
docAggs = append(docAggs, map[string]any{
"doc_name": title,
"doc_id": documentID,
"count": 1,
"url": resultURL,
})
}
return chunks, docAggs
}
// BuildComponentOutputs keeps the complete upstream object in json and adds a
// reference-formatted string for downstream Canvas components.
func (q *QueritTool) BuildComponentOutputs(response map[string]any) map[string]any {
chunks, _ := q.BuildReferences(context.Background(), response)
return map[string]any{
"formalized_content": renderQueritReferences(chunks, queritPromptMaxTokens),
"json": response,
}
}
func queritResults(response map[string]any) []map[string]any {
results, ok := response["results"].(map[string]any)
if !ok {
return nil
}
raw, ok := results["result"].([]any)
if !ok {
if typed, typedOK := results["result"].([]map[string]any); typedOK {
return typed
}
return nil
}
items := make([]map[string]any, 0, len(raw))
for _, value := range raw {
if item, itemOK := value.(map[string]any); itemOK {
items = append(items, item)
}
}
return items
}
func renderQueritReferences(chunks []map[string]any, maxTokens int) string {
usedTokens := 0
blocks := make([]string, 0, len(chunks))
for _, chunk := range chunks {
content := queritText(chunk["content"])
usedTokens += tokenizer.NumTokensFromString(content)
blocks = append(blocks, strings.Join([]string{
"\nID: " + queritText(chunk["id"]),
"├── Title: " + queritPromptField(chunk["document_name"]),
"├── URL: " + queritPromptField(chunk["url"]),
"└── Content:\n" + content,
}, "\n"))
if maxTokens > 0 && float64(maxTokens)*0.97 < float64(usedTokens) {
break
}
}
return strings.Join(blocks, "\n")
}
func queritText(value any) string {
if value == nil {
return ""
}
if text, ok := value.(string); ok {
return text
}
return fmt.Sprint(value)
}
func queritPromptField(value any) string {
return queritNewlinePattern.ReplaceAllString(queritText(value), " ")
}
func queritHashInt(value string, modulus int64) int64 {
digest := sha1.Sum([]byte(value))
number := new(big.Int).SetBytes(digest[:])
return new(big.Int).Mod(number, big.NewInt(modulus)).Int64()
}
func truncateQueritRunes(value string, limit int) string {
if limit <= 0 {
return ""
}
runes := []rune(value)
if len(runes) <= limit {
return value
}
return string(runes[:limit])
}
func queritInt(value int) *int { return &value }
func queritStringSlice(value any) ([]string, bool) {
switch items := value.(type) {
case []string:
return append([]string(nil), items...), true
case []any:
result := make([]string, 0, len(items))
for _, item := range items {
text, ok := item.(string)
if !ok {
return nil, false
}
result = append(result, text)
}
return result, true
case nil:
return nil, true
default:
return nil, false
}
}
func queritErrorJSON(err error, apiKeys ...string) string {
message := "querit_search: unknown error"
if err != nil {
message = "querit_search: " + err.Error()
}
for _, apiKey := range apiKeys {
if apiKey = strings.TrimSpace(apiKey); apiKey != "" {
message = strings.ReplaceAll(message, apiKey, "[REDACTED]")
}
}
raw, marshalErr := json.Marshal(map[string]any{"_ERROR": message})
if marshalErr != nil {
return `{"_ERROR":"querit_search: marshal error"}`
}
return string(raw)
}

View File

@@ -0,0 +1,594 @@
//
// 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 tool
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"sync/atomic"
"testing"
"time"
)
func TestQueritBuildsMinimalRequest(t *testing.T) {
var gotMethod, gotPath, gotAuthorization, gotContentType string
var gotBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
gotMethod = request.Method
gotPath = request.URL.Path
gotAuthorization = request.Header.Get("Authorization")
gotContentType = request.Header.Get("Content-Type")
_ = json.NewDecoder(request.Body).Decode(&gotBody)
writer.Header().Set("Content-Type", "application/json")
_, _ = writer.Write([]byte(`{"results":{"result":[]},"search_id":"search-1"}`))
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := newQueritTool(helper, func() string { return "" }, queritParams{APIKey: "key-test"}, nil)
out, err := querit.InvokableRun(context.Background(), `{"query":"ragflow"}`)
if err != nil {
t.Fatalf("InvokableRun: %v", err)
}
if gotMethod != http.MethodPost || gotPath != "/v1/search" {
t.Fatalf("request = %s %s, want POST /v1/search", gotMethod, gotPath)
}
if gotAuthorization != "Bearer key-test" {
t.Fatalf("Authorization = %q", gotAuthorization)
}
if !strings.HasPrefix(gotContentType, "application/json") {
t.Fatalf("Content-Type = %q", gotContentType)
}
if gotBody["query"] != "ragflow" || gotBody["count"] != float64(10) || gotBody["chunksPerDoc"] != float64(3) {
t.Fatalf("request body = %#v", gotBody)
}
if _, exists := gotBody["filters"]; exists {
t.Fatalf("empty filters must be omitted: %#v", gotBody)
}
if !strings.Contains(out, `"search_id":"search-1"`) {
t.Fatalf("complete response was not retained: %s", out)
}
}
func TestQueritBuildsFiltersAndMergesRuntimeOverrides(t *testing.T) {
var gotBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
_ = json.NewDecoder(request.Body).Decode(&gotBody)
writer.Header().Set("Content-Type", "application/json")
_, _ = writer.Write([]byte(`{"results":{"result":[]}}`))
}))
defer server.Close()
defaults := queritParams{
APIKey: "stored-key",
Count: 20,
ChunksPerDoc: queritInt(2),
SiteInclude: []string{"stored.example"},
SiteExclude: []string{"blocked.example"},
TimeRange: "w1",
CountryInclude: []string{"CN"},
LanguageInclude: []string{"zh"},
}
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := newQueritTool(helper, func() string { return "" }, defaults, func(context.Context, int) bool { return true })
_, err := querit.InvokableRun(context.Background(), `{"query":"ragflow","count":5,"site_include":[],"language_include":["en"]}`)
if err != nil {
t.Fatalf("InvokableRun: %v", err)
}
if gotBody["count"] != float64(5) || gotBody["chunksPerDoc"] != float64(2) {
t.Fatalf("merged scalar defaults = %#v", gotBody)
}
filters := gotBody["filters"].(map[string]any)
sites := filters["sites"].(map[string]any)
if _, exists := sites["include"]; exists || len(sites["exclude"].([]any)) != 1 {
t.Fatalf("explicit empty site_include did not clear node default: %#v", sites)
}
if filters["timeRange"].(map[string]any)["date"] != "w1" {
t.Fatalf("timeRange = %#v", filters["timeRange"])
}
if filters["geo"].(map[string]any)["countries"].(map[string]any)["include"].([]any)[0] != "CN" {
t.Fatalf("geo filter = %#v", filters["geo"])
}
if filters["languages"].(map[string]any)["include"].([]any)[0] != "en" {
t.Fatalf("language filter = %#v", filters["languages"])
}
}
func TestQueritUsesNodeQueryWhenRuntimeQueryIsOmitted(t *testing.T) {
var gotBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
_ = json.NewDecoder(request.Body).Decode(&gotBody)
writer.Header().Set("Content-Type", "application/json")
_, _ = writer.Write([]byte(`{"results":{"result":[]}}`))
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := newQueritTool(
helper,
func() string { return "" },
queritParams{APIKey: "stored-key", Query: "node query"},
nil,
)
out, err := querit.InvokableRun(context.Background(), `{}`)
if err != nil || strings.Contains(out, "_ERROR") {
t.Fatalf("InvokableRun = %s, %v", out, err)
}
if gotBody["query"] != "node query" {
t.Fatalf("request query = %#v, want node query", gotBody["query"])
}
}
func TestQueritAPIKeyResolutionAndEmptyQuery(t *testing.T) {
var calls atomic.Int32
var authorization string
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
calls.Add(1)
authorization = request.Header.Get("Authorization")
_, _ = writer.Write([]byte(`{"results":{"result":[]}}`))
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := NewQueritToolWithEnvKey(helper, func() string { return "environment-secret" })
if _, err := querit.InvokableRun(context.Background(), `{"query":"ragflow"}`); err != nil {
t.Fatalf("InvokableRun: %v", err)
}
if authorization != "Bearer environment-secret" {
t.Fatalf("Authorization = %q", authorization)
}
emptyOut, err := querit.InvokableRun(context.Background(), `{"query":""}`)
if err != nil {
t.Fatalf("empty query: %v", err)
}
if calls.Load() != 1 || emptyOut != `{}` {
t.Fatalf("empty query result = %s; calls = %d", emptyOut, calls.Load())
}
missing := NewQueritToolWithEnvKey(helper, func() string { return "" })
out, err := missing.InvokableRun(context.Background(), `{"query":"ragflow"}`)
if err != nil {
t.Fatalf("missing key returned Go error: %v", err)
}
if !strings.Contains(out, "api_key") || strings.Contains(out, "environment-secret") || calls.Load() != 1 {
t.Fatalf("missing-key result = %s, calls = %d", out, calls.Load())
}
}
func TestQueritRuntimeAPIKeyCannotOverrideNodeConfiguration(t *testing.T) {
var authorization string
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
authorization = request.Header.Get("Authorization")
_, _ = writer.Write([]byte(`{"results":{"result":[]}}`))
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := newQueritTool(helper, func() string { return "" }, queritParams{APIKey: "stored-key"}, nil)
out, err := querit.InvokableRun(context.Background(), `{"query":"ragflow","api_key":"runtime-key"}`)
if err != nil || strings.Contains(out, "_ERROR") {
t.Fatalf("InvokableRun = %s, %v", out, err)
}
if authorization != "Bearer stored-key" {
t.Fatalf("Authorization = %q, want stored node key", authorization)
}
}
func TestQueritRejectsMissingNullAndNonStringQueries(t *testing.T) {
var calls atomic.Int32
helper := NewHTTPHelper().WithClient(&http.Client{Transport: roundTripperErrorFunc(func(*http.Request) error {
calls.Add(1)
return errors.New("network must not be called")
})})
querit := NewQueritToolWith(helper)
for _, args := range []string{`{}`, `{"query":null}`, `{"query":123}`} {
t.Run(args, func(t *testing.T) {
out, err := querit.InvokableRun(context.Background(), args)
if err != nil || !strings.Contains(out, "_ERROR") || !strings.Contains(out, "query") {
t.Fatalf("InvokableRun(%s) = %s, %v", args, out, err)
}
})
}
if calls.Load() != 0 {
t.Fatalf("invalid queries made %d network calls", calls.Load())
}
}
func TestQueritExplicitNullChunksPerDocIsOmitted(t *testing.T) {
var gotBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
_ = json.NewDecoder(request.Body).Decode(&gotBody)
_, _ = writer.Write([]byte(`{"results":{"result":[]}}`))
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := NewQueritToolWithEnvKey(helper, func() string { return "k" })
out, err := querit.InvokableRun(context.Background(), `{"query":"x","chunks_per_doc":null}`)
if err != nil || strings.Contains(out, "_ERROR") {
t.Fatalf("InvokableRun = %s, %v", out, err)
}
if _, exists := gotBody["chunksPerDoc"]; exists {
t.Fatalf("explicit null chunks_per_doc was not omitted: %#v", gotBody)
}
}
func TestQueritValidatesParametersBeforeRequest(t *testing.T) {
tests := []string{
`{"query":"x","count":0}`,
`{"query":"x","chunks_per_doc":4}`,
`{"query":"x","time_range":"last week"}`,
`{"query":"x","time_range":"2026-01-01,2026-01-31"}`,
`{"query":"x","site_include":[1]}`,
}
for _, args := range tests {
t.Run(args, func(t *testing.T) {
out, err := NewQueritTool().InvokableRun(context.Background(), args)
if err != nil || !strings.Contains(out, "_ERROR") {
t.Fatalf("InvokableRun(%s) = %s, %v", args, out, err)
}
})
}
}
func TestQueritTimeRangeContract(t *testing.T) {
tests := []struct {
value string
valid bool
}{
{value: "", valid: true},
{value: "d7", valid: true},
{value: "w2", valid: true},
{value: "m3", valid: true},
{value: "y1", valid: true},
{value: "2026-01-01to2026-01-31", valid: true},
{value: "d0", valid: false},
{value: "7d", valid: false},
{value: "2026-01-01,2026-01-31", valid: false},
{value: "2026-01-01..2026-01-31", valid: false},
{value: "2026-01-01-2026-01-31", valid: false},
{value: "last week", valid: false},
}
for _, test := range tests {
t.Run(test.value, func(t *testing.T) {
if got := isValidQueritTimeRange(test.value); got != test.valid {
t.Fatalf("isValidQueritTimeRange(%q) = %v, want %v", test.value, got, test.valid)
}
if !test.valid || test.value == "" {
return
}
request := buildQueritRequest(queritParams{TimeRange: test.value})
if request.Filters == nil || request.Filters.TimeRange == nil || request.Filters.TimeRange.Date != test.value {
t.Fatalf("time range request mapping = %#v", request.Filters)
}
})
}
}
func TestQueritHTTPFailuresAreSoftErrors(t *testing.T) {
t.Run("unauthorized is not retried", func(t *testing.T) {
var calls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
calls.Add(1)
writer.WriteHeader(http.StatusUnauthorized)
_, _ = writer.Write([]byte(`{"message":"do not expose upstream bodies"}`))
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := NewQueritToolWithEnvKey(helper, func() string { return "secret-key" })
out, err := querit.InvokableRun(context.Background(), `{"query":"x"}`)
if err != nil || calls.Load() != 1 || !strings.Contains(out, "401") || strings.Contains(out, "secret-key") {
t.Fatalf("result = %s, err = %v, calls = %d", out, err, calls.Load())
}
})
t.Run("rate limit is retried at most three times", func(t *testing.T) {
var calls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
calls.Add(1)
writer.WriteHeader(http.StatusTooManyRequests)
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := newQueritTool(helper, nil, queritParams{APIKey: "secret-key"}, func(context.Context, int) bool { return true })
out, err := querit.InvokableRun(context.Background(), `{"query":"x"}`)
if err != nil || calls.Load() != queritMaxAttempts || !strings.Contains(out, "429") || strings.Contains(out, "secret-key") {
t.Fatalf("result = %s, err = %v, calls = %d", out, err, calls.Load())
}
})
t.Run("HTTP helper retries temporary server errors", func(t *testing.T) {
var calls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
attempt := calls.Add(1)
if attempt < 3 {
writer.WriteHeader(http.StatusServiceUnavailable)
return
}
_, _ = writer.Write([]byte(`{"results":{"result":[]}}`))
}))
defer server.Close()
helper := NewHTTPHelperWithRetry(RetryConfig{
MaxAttempts: 3,
BaseBackoff: time.Nanosecond,
MaxBackoff: time.Nanosecond,
}).WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
out, err := NewQueritToolWithEnvKey(helper, func() string { return "k" }).InvokableRun(context.Background(), `{"query":"x"}`)
if err != nil || calls.Load() != 3 || strings.Contains(out, "_ERROR") {
t.Fatalf("result = %s, err = %v, calls = %d", out, err, calls.Load())
}
})
t.Run("persistent server errors are soft", func(t *testing.T) {
var calls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
calls.Add(1)
writer.WriteHeader(http.StatusInternalServerError)
}))
defer server.Close()
helper := NewHTTPHelperWithRetry(RetryConfig{
MaxAttempts: 3,
BaseBackoff: time.Nanosecond,
MaxBackoff: time.Nanosecond,
}).WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := NewQueritToolWithEnvKey(helper, func() string { return "environment-secret" })
out, err := querit.InvokableRun(context.Background(), `{"query":"x"}`)
if err != nil || calls.Load() != 3 || !strings.Contains(out, "_ERROR") || !strings.Contains(out, "500") {
t.Fatalf("result = %s, err = %v, calls = %d", out, err, calls.Load())
}
})
for _, test := range []struct {
name string
secret string
node bool
}{
{name: "node API key is redacted from transport errors", secret: "node-secret", node: true},
{name: "environment API key is redacted from transport errors", secret: "environment-secret"},
} {
t.Run(test.name, func(t *testing.T) {
helper := NewHTTPHelperWithRetry(RetryConfig{
MaxAttempts: 1,
BaseBackoff: time.Nanosecond,
MaxBackoff: time.Nanosecond,
}).WithClient(&http.Client{Transport: roundTripperErrorFunc(func(*http.Request) error {
return fmt.Errorf("transport rejected Bearer %s", test.secret)
})})
querit := NewQueritToolWithEnvKey(helper, func() string { return test.secret })
if test.node {
querit = newQueritTool(helper, func() string { return "" }, queritParams{APIKey: test.secret}, nil)
}
out, err := querit.InvokableRun(context.Background(), `{"query":"x"}`)
if err != nil || !strings.Contains(out, "_ERROR") || !strings.Contains(out, "[REDACTED]") {
t.Fatalf("result = %s, err = %v", out, err)
}
if strings.Contains(out, test.secret) {
t.Fatalf("soft error exposed API key: %s", out)
}
})
}
t.Run("network errors are soft", func(t *testing.T) {
var calls atomic.Int32
helper := NewHTTPHelperWithRetry(RetryConfig{
MaxAttempts: 3,
BaseBackoff: time.Nanosecond,
MaxBackoff: time.Nanosecond,
}).WithClient(&http.Client{Transport: roundTripperErrorFunc(func(*http.Request) error {
calls.Add(1)
return errors.New("offline")
})})
out, err := NewQueritToolWithEnvKey(helper, func() string { return "k" }).InvokableRun(context.Background(), `{"query":"x"}`)
if err != nil || calls.Load() != 3 || !strings.Contains(out, "_ERROR") {
t.Fatalf("result = %s, err = %v, calls = %d", out, err, calls.Load())
}
})
}
func TestQueritRejectsInvalidJSONResponse(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte(`not-json`))
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := NewQueritToolWithEnvKey(helper, func() string { return "k" })
out, err := querit.InvokableRun(context.Background(), `{"query":"x"}`)
if err != nil || !strings.Contains(out, "decode response") {
t.Fatalf("result = %s, err = %v", out, err)
}
}
func TestQueritRejectsMalformedResponseShapes(t *testing.T) {
tests := []struct {
name string
body string
want string
}{
{name: "top-level array", body: `[]`, want: "JSON object"},
{name: "results is not an object", body: `{"results":[]}`, want: "results must be a JSON object"},
{name: "results is null", body: `{"results":null}`, want: "results must be a JSON object"},
{name: "result is not an array", body: `{"results":{"result":{}}}`, want: "results.result must be a JSON array"},
{name: "result is null", body: `{"results":{"result":null}}`, want: "results.result must be a JSON array"},
{name: "trailing content", body: `{"results":{"result":[]}} trailing`, want: "trailing content"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte(test.body))
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := NewQueritToolWithEnvKey(helper, func() string { return "k" })
out, err := querit.InvokableRun(context.Background(), `{"query":"x"}`)
if err != nil || !strings.Contains(out, "_ERROR") || !strings.Contains(out, test.want) {
t.Fatalf("result = %s, err = %v", out, err)
}
})
}
}
func TestQueritAcceptsMissingResultContainers(t *testing.T) {
for _, body := range []string{`{}`, `{"results":{}}`} {
t.Run(body, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte(body))
}))
defer server.Close()
helper := NewHTTPHelper().WithClient(&http.Client{Transport: rewriteQueritHostTransport(server.URL)})
querit := NewQueritToolWithEnvKey(helper, func() string { return "k" })
out, err := querit.InvokableRun(context.Background(), `{"query":"x"}`)
if err != nil || strings.Contains(out, "_ERROR") {
t.Fatalf("result = %s, err = %v", out, err)
}
})
}
}
func TestQueritReferencesAndCompleteComponentOutput(t *testing.T) {
response := map[string]any{
"search_id": "search-1",
"query_context": map[string]any{"rewritten": "rag flow"},
"results": map[string]any{"result": []any{
map[string]any{"title": "RAGFlow", "url": "https://ragflow.io", "snippet": "RAG engine", "extra": true},
map[string]any{"title": "Missing optional values"},
"invalid item",
}},
}
querit := NewQueritTool()
chunks, docAggs := querit.BuildReferences(context.Background(), response)
if len(chunks) != 1 || len(docAggs) != 1 {
t.Fatalf("references = %#v / %#v", chunks, docAggs)
}
if chunks[0]["content"] != "RAG engine" || chunks[0]["score"] != 1 || chunks[0]["similarity"] != 1 {
t.Fatalf("reference = %#v", chunks[0])
}
outputs := querit.BuildComponentOutputs(response)
if outputs["json"].(map[string]any)["search_id"] != "search-1" {
t.Fatalf("complete json output = %#v", outputs["json"])
}
formatted := outputs["formalized_content"].(string)
for _, expected := range []string{"Title: RAGFlow", "URL: https://ragflow.io", "RAG engine"} {
if !strings.Contains(formatted, expected) {
t.Fatalf("formalized_content missing %q: %s", expected, formatted)
}
}
if chunks, docAggs := querit.BuildReferences(context.Background(), map[string]any{"results": nil}); len(chunks) != 0 || len(docAggs) != 0 {
t.Fatalf("malformed response references = %#v / %#v", chunks, docAggs)
}
}
func TestQueritReferencesSanitizeAndLimitSnippets(t *testing.T) {
longSnippet := strings.Repeat("界", 10001)
response := map[string]any{"results": map[string]any{"result": []any{
map[string]any{"title": "empty", "snippet": ""},
map[string]any{"title": "image only", "snippet": "![img](data:image/png;base64,AAAA)"},
map[string]any{"title": "cleaned", "snippet": "before ![img](data:image/png;base64,AAAA) after"},
map[string]any{"title": "limited", "snippet": longSnippet},
}}}
chunks, docAggs := NewQueritTool().BuildReferences(context.Background(), response)
if len(chunks) != 2 || len(docAggs) != 2 {
t.Fatalf("references = %#v / %#v", chunks, docAggs)
}
if chunks[0]["content"] != "before after" {
t.Fatalf("base64 image was not removed: %#v", chunks[0])
}
limited, _ := chunks[1]["content"].(string)
if len([]rune(limited)) != 10000 {
t.Fatalf("limited snippet length = %d", len([]rune(limited)))
}
}
func TestQueritInfoDoesNotExposeAPIKey(t *testing.T) {
info, err := NewQueritTool().Info(context.Background())
if err != nil {
t.Fatalf("Info: %v", err)
}
if info.Name != queritToolName || info.Desc == "" || info.ParamsOneOf == nil {
t.Fatalf("Info = %#v", info)
}
encoded, err := json.Marshal(info)
if err != nil {
t.Fatalf("marshal Info: %v", err)
}
if strings.Contains(string(encoded), "api_key") {
t.Fatalf("Info exposed API key: %s", encoded)
}
jsonSchema, err := info.ParamsOneOf.ToJSONSchema()
if err != nil {
t.Fatalf("ToJSONSchema: %v", err)
}
rawSchema, err := json.Marshal(jsonSchema)
if err != nil {
t.Fatalf("marshal schema: %v", err)
}
var paramsSchema map[string]any
if err := json.Unmarshal(rawSchema, &paramsSchema); err != nil {
t.Fatalf("decode schema: %v", err)
}
properties, ok := paramsSchema["properties"].(map[string]any)
if !ok {
t.Fatalf("schema properties = %#v", paramsSchema["properties"])
}
for _, name := range []string{"site_include", "site_exclude", "country_include", "language_include"} {
property, ok := properties[name].(map[string]any)
if !ok {
t.Fatalf("%s schema = %#v", name, properties[name])
}
items, ok := property["items"].(map[string]any)
if !ok || items["type"] != "string" {
t.Fatalf("%s items schema = %#v", name, property["items"])
}
}
}
type roundTripperErrorFunc func(*http.Request) error
func (f roundTripperErrorFunc) RoundTrip(request *http.Request) (*http.Response, error) {
return nil, f(request)
}
func rewriteQueritHostTransport(serverURL string) http.RoundTripper {
target, err := url.Parse(serverURL)
if err != nil {
panic("rewriteQueritHostTransport: bad server URL: " + err.Error())
}
return &queritHostSwapTransport{
inner: http.DefaultTransport,
host: target.Host,
scheme: target.Scheme,
}
}
type queritHostSwapTransport struct {
inner http.RoundTripper
host string
scheme string
}
func (transport *queritHostSwapTransport) RoundTrip(request *http.Request) (*http.Response, error) {
cloned := request.Clone(request.Context())
cloned.URL.Scheme = transport.scheme
cloned.URL.Host = transport.host
cloned.Host = transport.host
return transport.inner.RoundTrip(cloned)
}

View File

@@ -49,6 +49,9 @@ var registry = map[string]Factory{
"keenable": buildKeenableTool,
"pubmed": buildPubMedTool,
"qweather": noConfig("qweather", func() einotool.BaseTool { return NewQWeatherTool() }),
"querit": buildQueritTool,
"querit_search": buildQueritTool,
"queritsearch": buildQueritTool,
"retrieval": buildRetrievalTool,
"search_my_dataset": buildRetrievalTool,
"search_my_dateset": buildRetrievalTool,
@@ -533,6 +536,61 @@ func buildTavilyTool(params map[string]any) (einotool.BaseTool, error) {
return newTavilyTool(nil, nil, defaults), nil
}
func buildQueritTool(params map[string]any) (einotool.BaseTool, error) {
defaults := queritParams{}
stringFields := map[string]*string{
"api_key": &defaults.APIKey,
"query": &defaults.Query,
"time_range": &defaults.TimeRange,
}
for key, destination := range stringFields {
value, exists := params[key]
if !exists {
continue
}
text, valid := value.(string)
if !valid {
return nil, fmt.Errorf("agent tool: tool %q requires string node-level param %s", "querit_search", key)
}
*destination = text
}
if value, exists := params["count"]; exists {
count, valid := strictInt(value)
if !valid || count < 1 {
return nil, fmt.Errorf("agent tool: tool %q requires integer node-level param count of at least 1", "querit_search")
}
defaults.Count = count
}
if value, exists := params["chunks_per_doc"]; exists {
chunksPerDoc, valid := strictInt(value)
if !valid || chunksPerDoc < 1 || chunksPerDoc > 3 {
return nil, fmt.Errorf("agent tool: tool %q requires integer node-level param chunks_per_doc within [1, 3]", "querit_search")
}
defaults.ChunksPerDoc = queritInt(chunksPerDoc)
}
listFields := map[string]*[]string{
"site_include": &defaults.SiteInclude,
"site_exclude": &defaults.SiteExclude,
"country_include": &defaults.CountryInclude,
"language_include": &defaults.LanguageInclude,
}
for key, destination := range listFields {
value, exists := params[key]
if !exists {
continue
}
items, valid := queritStringSlice(value)
if !valid {
return nil, fmt.Errorf("agent tool: tool %q requires string array node-level param %s", "querit_search", key)
}
*destination = items
}
if !isValidQueritTimeRange(strings.TrimSpace(defaults.TimeRange)) {
return nil, fmt.Errorf("agent tool: tool %q has invalid node-level param time_range", "querit_search")
}
return newQueritTool(nil, nil, defaults, nil), nil
}
func buildKeenableTool(params map[string]any) (einotool.BaseTool, error) {
defaults := keenableParams{}
apiKey := ""

View File

@@ -64,12 +64,43 @@ func TestBuildByName_TavilyCanvasComponentNames(t *testing.T) {
}
}
func TestBuildByName_QueritAliases(t *testing.T) {
for _, name := range []string{"querit", "querit_search", "queritsearch", "QueritSearch"} {
built, err := BuildByName(name, map[string]any{
"api_key": "stored-key",
"count": float64(8),
"chunks_per_doc": float64(2),
"site_include": []any{"example.com"},
"site_exclude": []string{"blocked.example"},
"time_range": "d7",
"country_include": []any{"CN"},
"language_include": []any{"zh"},
"outputs": map[string]any{"json": map[string]any{}},
})
if err != nil {
t.Fatalf("BuildByName(%q): %v", name, err)
}
querit, ok := built.(*QueritTool)
if !ok {
t.Fatalf("BuildByName(%q) returned %T, want *QueritTool", name, built)
}
if querit.defaults.Count != 8 || querit.defaults.ChunksPerDoc == nil || *querit.defaults.ChunksPerDoc != 2 || len(querit.defaults.SiteInclude) != 1 {
t.Fatalf("BuildByName(%q) defaults = %#v", name, querit.defaults)
}
info, infoErr := built.Info(context.Background())
if infoErr != nil || info.Name != "querit_search" {
t.Fatalf("BuildByName(%q).Info() = %#v, %v", name, info, infoErr)
}
}
}
func TestBuildAll_AllRegisteredTools(t *testing.T) {
// Every key in registry.
names := []string{
"akshare", "arxiv", "bgpt", "code_exec", "crawler", "deepl",
"duckduckgo", "email", "exesql", "execute_sql", "github", "google",
"google_scholar", "google_scholar_search", "jin10", "keenable", "pubmed", "qweather",
"querit", "querit_search", "queritsearch",
"retrieval", "search_my_dataset", "search_my_dateset", "searxng",
"tavily", "tavily_extract", "tushare", "web_crawler", "wencai", "wikipedia", "wikipedia_search",
"yahoo_finance",
@@ -133,6 +164,7 @@ func TestToolRegistry_SchemasAreComplete(t *testing.T) {
"akshare", "arxiv", "bgpt", "code_exec", "crawler", "deepl",
"duckduckgo", "email", "execute_sql", "exesql", "github", "google",
"google_scholar", "google_scholar_search", "jin10", "keenable", "pubmed", "qweather",
"querit", "querit_search", "queritsearch",
"retrieval", "search_my_dataset", "search_my_dateset", "searxng",
"tavily", "tavily_extract", "tushare", "web_crawler", "wencai", "wikipedia", "wikipedia_search",
"yahoo_finance",
@@ -202,6 +234,9 @@ func TestToolRegistry_SchemasAreComplete(t *testing.T) {
"web_crawler": "web_crawler",
"wikipedia": "wikipedia_search",
"wikipedia_search": "wikipedia_search",
"querit": "querit_search",
"querit_search": "querit_search",
"queritsearch": "querit_search",
}
for _, name := range names {
canonical, ok := canonicalByAlias[name]

View File

@@ -34,6 +34,9 @@ type ComponentSpec struct {
Inputs map[string]string
Outputs map[string]string
InputForm map[string]any
// PreserveJSONNumbers keeps numeric response tokens as json.Number when
// Canvas must expose the upstream JSON without float64 precision loss.
PreserveJSONNumbers bool
}
// ToolComponent is the required Canvas adaptation contract implemented by a

View File

@@ -113,6 +113,7 @@ const (
EnvSSHEnableAPIURL = "SSH_ENABLE_API_URL"
EnvAllowAnyHost = "ALLOW_ANY_HOST"
EnvTavilyApiKey = "TAVILY_API_KEY"
EnvQueritAPIKey = "QUERIT_API_KEY"
EnvHome = "HOME"
EnvUserProfile = "USERPROFILE"
EnvHttpHTTPProxy = "http_proxy"