package llm import ( "fmt" "slices" "testing" ) type byteTokenizer struct{} func (bt byteTokenizer) Tokenize(text string) []int { r := make([]int, len(text)) for i := 0; i < len(text); i++ { r[i] = int(text[i]) } return r } func TestTokenBuffer_Push(t *testing.T) { type gold struct { window int stride int text string expected [][]int64 } tests := []gold{ { window: 5, stride: 5, text: "abcdefghij", expected: [][]int64{{97, 98, 99, 100, 101}, {102, 103, 104, 105, 106}}, }, { window: 4, stride: 2, text: "abcdef", expected: [][]int64{{97, 98, 99, 100}, {99, 100, 101, 102}}, }, { window: 10, stride: 5, text: "abcde", expected: nil, }, } for _, tt := range tests { t.Run( fmt.Sprintf("window%d_stride%d", tt.window, tt.stride), func(t *testing.T) { tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: tt.window, Stride: tt.stride}) tb.SetIncludeTail(false) var got [][]int64 for w := range tb.Push(0, tt.text) { got = append(got, w.Tokens) } if len(got) != len(tt.expected) { t.Fatalf("expected %d windows but got %d", len(tt.expected), len(got)) } for i := range got { if !slices.Equal(got[i], tt.expected[i]) { t.Errorf("window %d: expected %v but got %v", i, tt.expected[i], got[i]) } } }, ) } } func TestTokenBuffer_PushAccumulates(t *testing.T) { tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 4, Stride: 4}) var got [][]int64 for w := range tb.Push(0, "ab") { got = append(got, w.Tokens) } if len(got) != 0 { t.Fatalf("expected 0 windows but got %d", len(got)) } for w := range tb.Push(0, "cd") { got = append(got, w.Tokens) } expected := [][]int64{{97, 98, 99, 100}} if len(got) != len(expected) { t.Fatalf("expected %d windows but got %d", len(expected), len(got)) } for i := range got { if !slices.Equal(got[i], expected[i]) { t.Errorf("window %d: expected %v but got %v", i, expected[i], got[i]) } } } func TestTokenBuffer_Tail(t *testing.T) { tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 5, Stride: 5}) tb.SetIncludeTail(true) var got [][]int64 for w := range tb.Push(0, "abcdef") { got = append(got, w.Tokens) } if w, ok := tb.Tail(); ok { got = append(got, w.Tokens) } expected := [][]int64{{97, 98, 99, 100, 101}, {102}} if len(got) != len(expected) { t.Fatalf("expected %d windows but got %d", len(expected), len(got)) } for i := range got { if !slices.Equal(got[i], expected[i]) { t.Errorf("window %d: expected %v but got %v", i, expected[i], got[i]) } } } func TestTokenBuffer_Position(t *testing.T) { tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 4, Stride: 2}) tb.SetIncludeTail(false) var positions []int for range tb.Push(0, "abcdefgh") { positions = append(positions, tb.Position()) } expected := []int{0, 2, 4} if len(positions) != len(expected) { t.Fatalf("expected %d positions but got %d", len(expected), len(positions)) } for i := range positions { if positions[i] != expected[i] { t.Errorf("window %d: expected positions=%d but got positions=%d", i, expected[i], positions[i]) } } tail, _ := tb.Tail() if tailPosition := positions[len(positions)-1] + tail.Seen; tailPosition != 6 { t.Errorf("expected tail positions=6 but got %d", tailPosition) } } func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) { tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 10, Stride: 10}) tb.SetIncludeTail(true) var a [][]int64 for w := range tb.Push(0, "abcde") { a = append(a, w.Tokens) } if len(a) != 0 { t.Fatalf("expected 0 windows but got %d", len(a)) } var b [][]int64 for w := range tb.Push(2, "fghij") { b = append(b, w.Tokens) } if len(b) != 1 { t.Fatalf("expected 1 window but got %d", len(b)) } expected := toInt64(byteTokenizer{}.Tokenize("abcde")) if !slices.Equal(b[0], expected) { t.Errorf("expected %v but got %v", expected, b[0]) } } func TestTokenBuffer_TailPosition(t *testing.T) { tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 4, Stride: 2}) tb.SetIncludeTail(true) var lastPosition int for range tb.Push(0, "abcdefgh") { lastPosition = tb.Position() } if lastPosition != 4 { t.Errorf("expected last position=4 but got %d", lastPosition) } tail, _ := tb.Tail() tailPosition := lastPosition + tail.Seen if !slices.Equal(tail.Tokens, toInt64(byteTokenizer{}.Tokenize("gh"))) { t.Errorf("unexpected tail: %v", tail.Tokens) } if tailPosition != 6 { t.Errorf("expected tail position=6 but got %d", tailPosition) } } func TestTokenBuffer_Pad(t *testing.T) { type gold struct { config TokenBufferConfig text string expected [][]int64 seen int } tests := []gold{ { config: TokenBufferConfig{ Window: 4, Stride: 2, PadLeft: true, PadRight: false, PadTokenID: 0, }, text: "abcdefg", expected: [][]int64{{97, 98, 99, 100}, {99, 100, 101, 102}, {0, 101, 102, 103}}, seen: 3, }, { config: TokenBufferConfig{ Window: 4, Stride: 2, PadLeft: false, PadRight: true, PadTokenID: 0, }, text: "abcdefg", expected: [][]int64{{97, 98, 99, 100}, {99, 100, 101, 102}, {101, 102, 103, 0}}, seen: 2, }, } for _, tt := range tests { t.Run( fmt.Sprintf("window%d_stride%d_%v_%v", tt.config.Window, tt.config.Stride, tt.config.PadLeft, tt.config.PadRight), func(t *testing.T) { tb := NewTokenBuffer(byteTokenizer{}, tt.config) tb.SetIncludeTail(true) var got [][]int64 for w := range tb.Push(0, tt.text) { got = append(got, w.Tokens) } tail, _ := tb.Tail() gotSeen := tail.Seen if tt.config.PadLeft { gotSeen += tail.PaddingLeft } if gotSeen != tt.seen { t.Fatalf("expected seen %d but got %d", tt.seen, gotSeen) } got = append(got, tail.Tokens) if len(got) != len(tt.expected) { t.Fatalf("expected %d windows but got %d", len(tt.expected), len(got)) } for i := range got { if !slices.Equal(got[i], tt.expected[i]) { t.Errorf("window %d: expected %v but got %v", i, tt.expected[i], got[i]) } } }, ) } }