Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 33 additions & 24 deletions internal/agent/tools/knowledge_search.go
Original file line number Diff line number Diff line change
Expand Up @@ -707,9 +707,16 @@ func (t *KnowledgeSearchTool) rerankWithLLM(
// This prevents token overflow and improves processing efficiency
const batchSize = 15
const maxContentLength = 800 // Maximum characters per passage to avoid excessive tokens
const reasoningTokenReserve = 1024

// Process in batches
allScores := make([]float64, len(results))
useOriginalScores := func(start, end int) {
for i := start; i < end; i++ {
allScores[i] = results[i].Score
}
}
disableThinking := false

for batchStart := 0; batchStart < len(results); batchStart += batchSize {
batchEnd := batchStart + batchSize
Expand Down Expand Up @@ -793,24 +800,39 @@ Output only the scores, no explanations or additional text.`,
},
}

// Calculate appropriate max tokens based on batch size
// Each score line is ~15 tokens, add buffer for safety
maxTokens := len(batch)*20 + 100
// Each score line is ~15 tokens. Disable reasoning for this structured
// scoring task and keep a reserve for providers that ignore that option.
maxTokens := len(batch)*20 + 100 + reasoningTokenReserve

modelCtx := types.WithLLMCallMetadata(ctx, "knowledge_search_rerank", "")
response, err := t.chatModel.Chat(modelCtx, messages, &chat.ChatOptions{
Temperature: 0.1, // Low temperature for consistent scoring
MaxTokens: maxTokens,
Thinking: &disableThinking,
})
if err != nil {
logger.Warnf(ctx, "[Tool][KnowledgeSearch] LLM rerank batch %d-%d failed: %v, using original scores",
batchStart+1, batchEnd, err)
// Use original scores for this batch on error
for i := batchStart; i < batchEnd; i++ {
allScores[i] = results[i].Score
}
useOriginalScores(batchStart, batchEnd)
continue
}
if response == nil || strings.TrimSpace(response.Content) == "" ||
strings.EqualFold(response.FinishReason, "length") {
finishReason := ""
if response != nil {
finishReason = response.FinishReason
}
logger.Warnf(
ctx,
"[Tool][KnowledgeSearch] LLM rerank batch %d-%d returned incomplete output (finish_reason=%q), using original scores for remaining results",
batchStart+1,
batchEnd,
finishReason,
)
useOriginalScores(batchStart, len(results))
break
}

logger.Infof(ctx, "[Tool][KnowledgeSearch] LLM rerank batch %d-%d response: %s",
batchStart+1, batchEnd, response.Content)
Expand All @@ -820,16 +842,13 @@ Output only the scores, no explanations or additional text.`,
if err != nil {
logger.Warnf(
ctx,
"[Tool][KnowledgeSearch] Failed to parse LLM scores for batch %d-%d: %v, using original scores",
"[Tool][KnowledgeSearch] Failed to parse LLM scores for batch %d-%d: %v, using original scores for remaining results",
batchStart+1,
batchEnd,
err,
)
// Use original scores for this batch on parsing error
for i := batchStart; i < batchEnd; i++ {
allScores[i] = results[i].Score
}
continue
useOriginalScores(batchStart, len(results))
break
}

// Store scores for this batch
Expand Down Expand Up @@ -919,18 +938,8 @@ func (t *KnowledgeSearchTool) parseScoresFromResponse(responseText string, expec
return nil, fmt.Errorf("no valid scores found in response")
}

// If we got fewer scores than expected, pad with last score or 0.5
for len(scores) < expectedCount {
if len(scores) > 0 {
scores = append(scores, scores[len(scores)-1])
} else {
scores = append(scores, 0.5)
}
}

// Truncate if we got more scores than expected
if len(scores) > expectedCount {
scores = scores[:expectedCount]
if len(scores) != expectedCount {
return nil, fmt.Errorf("expected %d scores, got %d", expectedCount, len(scores))
}

return scores, nil
Expand Down
131 changes: 131 additions & 0 deletions internal/agent/tools/knowledge_search_rerank_test.go
Original file line number Diff line number Diff line change
@@ -1,13 +1,68 @@
package tools

import (
"context"
"testing"

"github.com/Tencent/WeKnora/internal/config"
"github.com/Tencent/WeKnora/internal/models/chat"
"github.com/Tencent/WeKnora/internal/models/rerank"
"github.com/Tencent/WeKnora/internal/types"
)

type rerankChatStub struct {
response *types.ChatResponse
calls int
options []*chat.ChatOptions
}

func (s *rerankChatStub) Chat(
_ context.Context,
_ []chat.Message,
opts *chat.ChatOptions,
) (*types.ChatResponse, error) {
s.calls++
copied := *opts
s.options = append(s.options, &copied)
return s.response, nil
}

func (*rerankChatStub) ChatStream(
context.Context,
[]chat.Message,
*chat.ChatOptions,
) (<-chan types.StreamResponse, error) {
stream := make(chan types.StreamResponse)
close(stream)
return stream, nil
}

func (*rerankChatStub) GetModelName() string { return "rerank-test" }

func (*rerankChatStub) GetModelID() string { return "rerank-test" }

func TestRerankChatStubStreamIsClosed(t *testing.T) {
t.Parallel()

stream, err := (&rerankChatStub{}).ChatStream(
context.Background(),
nil,
nil,
)
if err != nil {
t.Fatalf("ChatStream() error = %v", err)
}

select {
case _, ok := <-stream:
if ok {
t.Fatal("ChatStream() returned an open stream")
}
default:
t.Fatal("ChatStream() stream is not closed")
}
}

func TestFilterRerankRankResults_thresholdAndFallback(t *testing.T) {
t.Parallel()
rankResults := []rerank.RankResult{
Expand Down Expand Up @@ -84,3 +139,79 @@ func TestRerankThreshold_default(t *testing.T) {
t.Fatalf("default threshold = %v, want 0.3", got)
}
}

func TestRerankWithLLMStopsAfterIncompleteOutput(t *testing.T) {
t.Parallel()

tests := []struct {
name string
response *types.ChatResponse
}{
{
name: "reasoning budget exhausted",
response: &types.ChatResponse{
FinishReason: "length",
},
},
{
name: "score list is incomplete",
response: &types.ChatResponse{
Content: "Passage 1: 0.90",
FinishReason: "stop",
},
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
stub := &rerankChatStub{response: tt.response}
tool := &KnowledgeSearchTool{chatModel: stub}
results := make([]*searchResultWithMeta, 16)
for i := range results {
results[i] = &searchResultWithMeta{SearchResult: &types.SearchResult{
ID: string(rune('a' + i)),
Content: "passage",
Score: 0.9,
}}
}

got, err := tool.rerankWithLLM(context.Background(), "query", results)
if err != nil {
t.Fatalf("rerankWithLLM returned error: %v", err)
}
if stub.calls != 1 {
t.Fatalf("chat calls = %d, want 1; invalid first batch should skip remaining batches", stub.calls)
}
if len(stub.options) != 1 {
t.Fatalf("captured options = %d, want 1", len(stub.options))
}
if stub.options[0].Thinking == nil || *stub.options[0].Thinking {
t.Fatalf("Thinking = %v, want false", stub.options[0].Thinking)
}
if stub.options[0].MaxTokens != 1424 {
t.Fatalf("MaxTokens = %d, want 1424", stub.options[0].MaxTokens)
}
if len(got) != len(results) {
t.Fatalf("reranked results = %d, want %d original-score fallbacks", len(got), len(results))
}
})
}
}

func TestParseScoresFromResponseRequiresExactCount(t *testing.T) {
t.Parallel()
tool := &KnowledgeSearchTool{}

scores, err := tool.parseScoresFromResponse("Passage 1: 0.90\nPassage 2: 0.40", 2)
if err != nil {
t.Fatalf("complete score list returned error: %v", err)
}
if len(scores) != 2 || scores[0] != 0.9 || scores[1] != 0.4 {
t.Fatalf("scores = %#v, want [0.9 0.4]", scores)
}

if _, err := tool.parseScoresFromResponse("Passage 1: 0.90", 2); err == nil {
t.Fatal("incomplete score list should return an error")
}
}
135 changes: 135 additions & 0 deletions internal/application/service/knowledge_batch_reparse_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,40 @@ type reparseFailureKBService struct {
kb *types.KnowledgeBase
}

type parserRulesKnowledgeRepo struct {
interfaces.KnowledgeRepository
knowledge *types.Knowledge
getErr error
updateErr error
requestedTenant uint64
updateCalls int
updatedID string
updatedColumn string
updatedValue interface{}
}

func (r *parserRulesKnowledgeRepo) GetKnowledgeByID(
_ context.Context,
tenantID uint64,
_ string,
) (*types.Knowledge, error) {
r.requestedTenant = tenantID
return r.knowledge, r.getErr
}

func (r *parserRulesKnowledgeRepo) UpdateKnowledgeColumn(
_ context.Context,
id string,
column string,
value interface{},
) error {
r.updateCalls++
r.updatedID = id
r.updatedColumn = column
r.updatedValue = value
return r.updateErr
}

func (s *reparseFailureKBService) GetKnowledgeBaseByID(
_ context.Context,
_ string,
Expand Down Expand Up @@ -134,3 +168,104 @@ func TestRunKnowledgeListReparseSubmissionsSucceeds(t *testing.T) {
require.NoError(t, err)
require.Equal(t, knowledgeListReparseOutcome{Submitted: 2}, outcome)
}

func TestClearStoredParserEngineRulesUsesCurrentKnowledgeBaseRules(t *testing.T) {
enableMultimodel := true
knowledge := &types.Knowledge{
ID: "knowledge-1",
TenantID: 7,
Metadata: types.JSON(`{
"source_id":"keep-me",
"process_overrides":{
"parser_engine_rules":[{"file_types":["pdf"],"engine":"old-top-level"}],
"chunking_config":{
"chunk_size":1024,
"chunk_overlap":128,
"enable_parent_child":true,
"parser_engine_rules":[{"file_types":["docx"],"engine":"old-nested"}]
},
"enable_multimodel":true,
"parser_engine_overrides":{"pdf_force_scanned":"true"}
}
}`),
}
repo := &parserRulesKnowledgeRepo{knowledge: knowledge}
svc := &knowledgeService{repo: repo}
ctx := context.WithValue(context.Background(), types.TenantIDContextKey, uint64(7))

err := svc.clearStoredParserEngineRules(ctx, knowledge.ID)

require.NoError(t, err)
require.Equal(t, uint64(7), repo.requestedTenant)
require.Equal(t, 1, repo.updateCalls)
require.Equal(t, knowledge.ID, repo.updatedID)
require.Equal(t, "metadata", repo.updatedColumn)
require.Equal(t, knowledge.Metadata, repo.updatedValue)

overrides, err := knowledge.ProcessOverrides()
require.NoError(t, err)
require.NotNil(t, overrides)
require.Empty(t, overrides.ParserEngineRules)
require.NotNil(t, overrides.ChunkingConfig)
require.Empty(t, overrides.ChunkingConfig.ParserEngineRules)
require.Equal(t, 1024, overrides.ChunkingConfig.ChunkSize)
require.Equal(t, 128, overrides.ChunkingConfig.ChunkOverlap)
require.True(t, overrides.ChunkingConfig.EnableParentChild)
require.Equal(t, &enableMultimodel, overrides.EnableMultimodel)
require.Equal(t, map[string]string{"pdf_force_scanned": "true"}, overrides.ParserEngineOverrides)

metadata, err := knowledge.Metadata.Map()
require.NoError(t, err)
require.Equal(t, "keep-me", metadata["source_id"])

currentRules := []types.ParserEngineRule{{FileTypes: []string{"pdf", "docx"}, Engine: "current"}}
effective := ResolveProcessConfig(&types.KnowledgeBase{
ChunkingConfig: types.ChunkingConfig{ParserEngineRules: currentRules},
}, overrides)
require.Equal(t, currentRules, effective.ChunkingConfig.ParserEngineRules)
}

func TestClearStoredParserEngineRulesWithoutSnapshotsIsNoOp(t *testing.T) {
knowledge := &types.Knowledge{
ID: "knowledge-2",
TenantID: 7,
Metadata: types.JSON(`{
"source_id":"keep-me",
"process_overrides":{
"chunking_config":{"chunk_size":2048},
"parser_engine_overrides":{"xlsx_first_row_as_header":"true"}
}
}`),
}
originalMetadata := append(types.JSON(nil), knowledge.Metadata...)
repo := &parserRulesKnowledgeRepo{knowledge: knowledge}
svc := &knowledgeService{repo: repo}
ctx := context.WithValue(context.Background(), types.TenantIDContextKey, uint64(7))

err := svc.clearStoredParserEngineRules(ctx, knowledge.ID)

require.NoError(t, err)
require.Zero(t, repo.updateCalls)
require.Equal(t, originalMetadata, knowledge.Metadata)
}

func TestClearStoredParserEngineRulesPropagatesUpdateFailure(t *testing.T) {
updateErr := errors.New("metadata update failed")
knowledge := &types.Knowledge{
ID: "knowledge-3",
TenantID: 7,
Metadata: types.JSON(`{
"process_overrides":{
"parser_engine_rules":[{"file_types":["pdf"],"engine":"old"}]
}
}`),
}
repo := &parserRulesKnowledgeRepo{knowledge: knowledge, updateErr: updateErr}
svc := &knowledgeService{repo: repo}
ctx := context.WithValue(context.Background(), types.TenantIDContextKey, uint64(7))

err := svc.clearStoredParserEngineRules(ctx, knowledge.ID)

require.ErrorIs(t, err, updateErr)
require.Equal(t, 1, repo.updateCalls)
}
Loading