mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-01 21:37:33 +08:00
595 lines
23 KiB
Go
595 lines
23 KiB
Go
//
|
|
// 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": ""},
|
|
map[string]any{"title": "cleaned", "snippet": "before  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, ¶msSchema); 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)
|
|
}
|