Files

298 lines
9.7 KiB
Go
Raw Permalink Normal View History

package messagefit
import (
"slices"
"strings"
"testing"
"ragflow/internal/tokenizer"
)
func TestFit_AllFits(t *testing.T) {
msgs := []Message{
{Role: "system", Content: "hello"},
{Role: "user", Content: "world"},
}
kept, keptIdx, count := Fit(msgs, 100000)
if count == 0 {
t.Errorf("Fit returned count 0, want > 0")
}
if len(kept) != 2 || !slices.Equal(keptIdx, []int{0, 1}) {
t.Fatalf("got kept=%+v keptIdx=%v, want both messages", kept, keptIdx)
}
if kept[0].Content != "hello" || kept[1].Content != "world" {
t.Errorf("messages modified when they fit: %+v", kept)
}
}
// TestFit_DoesNotMutateInput verifies the non-mutating contract: Fit never
// rewrites msgs, so a caller cannot accidentally send emptied entries.
func TestFit_DoesNotMutateInput(t *testing.T) {
orig := []Message{
{Role: "system", Content: strings.Repeat("s ", 500)},
{Role: "user", Content: strings.Repeat("u ", 500)},
{Role: "user", Content: "last"},
}
msgs := slices.Clone(orig)
Fit(msgs, 100)
if !slices.Equal(msgs, orig) {
t.Fatalf("Fit mutated its input:\n got %+v\nwant %+v", msgs, orig)
}
}
func TestFit_Step2_DropsMiddle(t *testing.T) {
// system + last user fit within the budget, but all three together do
// not, so Step 2 drops the middle user and keeps system + last intact.
sysContent := strings.Repeat("x ", 200)
middle := "middle"
last := "last"
msgs := []Message{
{Role: "system", Content: sysContent},
{Role: "user", Content: middle},
{Role: "user", Content: last},
}
budget := tokenizer.NumTokensFromString(sysContent) + tokenizer.NumTokensFromString(last)
kept, keptIdx, count := Fit(msgs, budget)
if count == 0 {
t.Fatalf("Fit returned count 0, want > 0")
}
if !slices.Equal(keptIdx, []int{0, 2}) {
t.Fatalf("keptIdx = %v, want [0 2] (middle dropped)", keptIdx)
}
if kept[0].Content != sysContent || kept[1].Content != last {
t.Errorf("retained messages modified: %+v", kept)
}
}
func TestFit_ExactBudget(t *testing.T) {
// A total exactly equal to the budget counts as fitting.
msgs := []Message{
{Role: "system", Content: "abc"},
{Role: "user", Content: "def"},
}
budget := tokenizer.NumTokensFromString("abc") + tokenizer.NumTokensFromString("def")
kept, keptIdx, _ := Fit(msgs, budget)
if len(kept) != 2 || !slices.Equal(keptIdx, []int{0, 1}) {
t.Fatalf("got kept=%+v keptIdx=%v, want both messages kept", kept, keptIdx)
}
if kept[0].Content != "abc" || kept[1].Content != "def" {
t.Errorf("messages modified at exact budget: %+v", kept)
}
}
// TestFit_NoSystem_KeepsOnlyLast locks the Python-parity behavior: without a
// system message, Step 2 keeps only the last non-system message.
func TestFit_NoSystem_KeepsOnlyLast(t *testing.T) {
msgs := []Message{
{Role: "user", Content: strings.Repeat("a ", 300)},
{Role: "assistant", Content: strings.Repeat("b ", 300)},
{Role: "user", Content: "last"},
}
kept, keptIdx, _ := Fit(msgs, 100)
if !slices.Equal(keptIdx, []int{2}) {
t.Fatalf("keptIdx = %v, want [2] (only the last user kept)", keptIdx)
}
if len(kept) != 1 || kept[0].Content != "last" {
t.Fatalf("kept = %+v, want the last user message", kept)
}
}
func TestFit_SystemOnlyMessages(t *testing.T) {
sys1 := strings.Repeat("s ", 800)
sys2 := strings.Repeat("t ", 200)
msgs := []Message{
{Role: "system", Content: sys1},
{Role: "system", Content: sys2},
}
const budget = 500
kept, keptIdx, count := Fit(msgs, budget)
if count == 0 {
t.Fatalf("Fit returned count 0, want > 0")
}
if len(kept) != 2 || !slices.Equal(keptIdx, []int{0, 1}) {
t.Fatalf("got kept=%+v keptIdx=%v, want both systems", kept, keptIdx)
}
total := tokenizer.NumTokensFromString(kept[0].Content) + tokenizer.NumTokensFromString(kept[1].Content)
if total > budget {
t.Errorf("fitted total %d exceeds budget %d", total, budget)
}
fitted0 := tokenizer.NumTokensFromString(kept[0].Content)
fitted1 := tokenizer.NumTokensFromString(kept[1].Content)
if fitted0 == 0 || fitted1 == 0 {
t.Errorf("a system message was emptied by fitting: %+v", kept)
}
if fitted0 >= tokenizer.NumTokensFromString(sys1) || fitted1 >= tokenizer.NumTokensFromString(sys2) {
t.Errorf("system messages not trimmed: %+v", kept)
}
}
func TestFit_TrimsAllSystemMessages(t *testing.T) {
sys1 := strings.Repeat("s ", 800)
sys2 := strings.Repeat("t ", 200)
last := "last"
msgs := []Message{
{Role: "system", Content: sys1},
{Role: "system", Content: sys2},
{Role: "user", Content: last},
}
const budget = 500
kept, keptIdx, count := Fit(msgs, budget)
if count == 0 {
t.Fatalf("Fit returned count 0, want > 0")
}
if !slices.Equal(keptIdx, []int{0, 1, 2}) {
t.Fatalf("keptIdx = %v, want [0 1 2]", keptIdx)
}
if kept[2].Content != last {
t.Errorf("last user message not preserved: %q", kept[2].Content)
}
total := tokenizer.NumTokensFromString(kept[0].Content) +
tokenizer.NumTokensFromString(kept[1].Content) +
tokenizer.NumTokensFromString(kept[2].Content)
if total > budget {
t.Errorf("fitted total %d exceeds budget %d", total, budget)
}
fitted0 := tokenizer.NumTokensFromString(kept[0].Content)
fitted1 := tokenizer.NumTokensFromString(kept[1].Content)
if fitted0 == 0 || fitted1 == 0 {
t.Errorf("a system message was emptied by fitting: %+v", kept)
}
if fitted0 >= tokenizer.NumTokensFromString(sys1) || fitted1 >= tokenizer.NumTokensFromString(sys2) {
t.Errorf("system messages not trimmed: %+v", kept)
}
}
func TestFit_Step3_SystemDominates(t *testing.T) {
// System takes >80% of tokens → preserve user, trim system.
sysContent := strings.Repeat("a ", 800) // dominates
userContent := strings.Repeat("b ", 100) // small, fits entirely
msgs := []Message{
{Role: "system", Content: sysContent},
{Role: "user", Content: userContent},
}
const budget = 500
kept, keptIdx, count := Fit(msgs, budget)
if count == 0 {
t.Fatalf("Fit returned count 0, want > 0")
}
if count > budget {
t.Errorf("fitted total %d exceeds budget %d", count, budget)
}
if !slices.Equal(keptIdx, []int{0, 1}) {
t.Fatalf("keptIdx = %v, want [0 1]", keptIdx)
}
// User preserved verbatim; system trimmed.
if tokenizer.NumTokensFromString(kept[1].Content) != tokenizer.NumTokensFromString(userContent) {
t.Errorf("user message not preserved: %+v", kept)
}
if tokenizer.NumTokensFromString(kept[0].Content) >= tokenizer.NumTokensFromString(sysContent) {
t.Errorf("system not trimmed: %+v", kept)
}
}
func TestFit_Step3_UserDominates(t *testing.T) {
// User takes >20% → preserve system, trim user.
sysContent := strings.Repeat("a ", 150) // small, fits entirely
userContent := strings.Repeat("b ", 800) // dominates
msgs := []Message{
{Role: "system", Content: sysContent},
{Role: "user", Content: userContent},
}
const budget = 500
kept, keptIdx, count := Fit(msgs, budget)
if count == 0 {
t.Fatalf("Fit returned count 0, want > 0")
}
if count > budget {
t.Errorf("fitted total %d exceeds budget %d", count, budget)
}
if !slices.Equal(keptIdx, []int{0, 1}) {
t.Fatalf("keptIdx = %v, want [0 1]", keptIdx)
}
// System preserved verbatim; user trimmed.
if tokenizer.NumTokensFromString(kept[0].Content) != tokenizer.NumTokensFromString(sysContent) {
t.Errorf("system message not preserved: %+v", kept)
}
if tokenizer.NumTokensFromString(kept[1].Content) >= tokenizer.NumTokensFromString(userContent) {
t.Errorf("user not trimmed: %+v", kept)
}
}
// TestFit_Step3_BudgetFilledByLast locks the boundary where the last message
// alone fills the budget (preserved == budget): the system share collapses to
// 0, so every retained system message is kept but trimmed to empty — it must
// NOT be reported as dropped. Python's message_fit_in retains the system entry
// with content "" in this case too.
func TestFit_Step3_BudgetFilledByLast(t *testing.T) {
sysContent := strings.Repeat("a ", 3000) // dominates (>80% of tokens)
userContent := strings.Repeat("b ", 600) // alone exceeds the budget
msgs := []Message{
{Role: "system", Content: sysContent},
{Role: "user", Content: userContent},
}
const budget = 500
kept, keptIdx, count := Fit(msgs, budget)
if !slices.Equal(keptIdx, []int{0, 1}) {
t.Fatalf("keptIdx = %v, want [0 1] (both messages still kept)", keptIdx)
}
if count > budget {
t.Errorf("fitted total %d exceeds budget %d", count, budget)
}
if kept[0].Content != "" {
t.Errorf("system message not trimmed to empty when the last message fills the budget: %+v", kept)
}
if tokenizer.NumTokensFromString(kept[1].Content) > budget {
t.Errorf("user message exceeds budget after trim: %+v", kept)
}
}
func TestFit_SingleMessage(t *testing.T) {
msgs := []Message{
{Role: "system", Content: strings.Repeat("x ", 1000)},
}
kept, keptIdx, count := Fit(msgs, 100)
if count == 0 {
t.Fatalf("Fit returned count 0, want > 0")
}
if !slices.Equal(keptIdx, []int{0}) {
t.Fatalf("keptIdx = %v, want [0]", keptIdx)
}
if tokenizer.NumTokensFromString(kept[0].Content) > 100 {
t.Errorf("single message not trimmed: %+v", kept)
}
}
func TestFit_ZeroBudget(t *testing.T) {
// budget <= 0 should use 8192 default and not panic.
msgs := []Message{
{Role: "system", Content: "hello"},
{Role: "user", Content: "world"},
}
kept, keptIdx, count := Fit(msgs, 0)
if count == 0 {
t.Fatalf("Fit returned count 0, want > 0")
}
if len(kept) != 2 || !slices.Equal(keptIdx, []int{0, 1}) {
t.Fatalf("got kept=%+v keptIdx=%v, want both messages", kept, keptIdx)
}
if kept[0].Content != "hello" || kept[1].Content != "world" {
t.Errorf("messages modified with default budget: %+v", kept)
}
}
func TestFit_Empty(t *testing.T) {
if kept, keptIdx, count := Fit(nil, 1000); kept != nil || keptIdx != nil || count != 0 {
t.Errorf("Fit(nil) = %v, %v, %d; want nil, nil, 0", kept, keptIdx, count)
}
if kept, keptIdx, count := Fit([]Message{}, 1000); len(kept) != 0 || len(keptIdx) != 0 || count != 0 {
t.Errorf("Fit(empty) = %v, %v, %d; want 0, 0, 0", kept, keptIdx, count)
}
}