88 lines
3.0 KiB
Go
88 lines
3.0 KiB
Go
package contextwindow
|
|
|
|
import (
|
|
"strings"
|
|
"testing"
|
|
|
|
"aichat/message"
|
|
|
|
"github.com/volcengine/volcengine-go-sdk/service/arkruntime/model"
|
|
)
|
|
|
|
func TestApplyChatWindowKeepsMessagesUnderLimit(t *testing.T) {
|
|
messages := []message.ChatMessage{{Role: "user", Content: "hello"}, {Role: "assistant", Content: "hi"}}
|
|
result := ApplyChatWindow(messages, 1000)
|
|
if result.Removed != 0 || result.Overflow {
|
|
t.Fatalf("unexpected result: %#v", result)
|
|
}
|
|
if len(result.ChatMessages) != len(messages) {
|
|
t.Fatalf("messages length = %d", len(result.ChatMessages))
|
|
}
|
|
}
|
|
|
|
func TestApplyChatWindowPrunesOldDialogueAndKeepsSystem(t *testing.T) {
|
|
messages := []message.ChatMessage{
|
|
{Role: "system", Content: strings.Repeat("规", 8)},
|
|
{Role: "user", Content: strings.Repeat("旧", 40)},
|
|
{Role: "assistant", Content: strings.Repeat("旧", 40)},
|
|
{Role: "user", Content: "最新问题"},
|
|
}
|
|
result := ApplyChatWindow(messages, 30)
|
|
if result.Removed == 0 {
|
|
t.Fatalf("expected old dialogue to be removed: %#v", result)
|
|
}
|
|
if len(result.ChatMessages) < 2 {
|
|
t.Fatalf("unexpected messages: %#v", result.ChatMessages)
|
|
}
|
|
if result.ChatMessages[0].Role != "system" {
|
|
t.Fatalf("system message not preserved: %#v", result.ChatMessages)
|
|
}
|
|
last := result.ChatMessages[len(result.ChatMessages)-1]
|
|
if last.Role != "user" || last.Content != "最新问题" {
|
|
t.Fatalf("latest user message not preserved: %#v", result.ChatMessages)
|
|
}
|
|
}
|
|
|
|
func TestApplyChatWindowReportsBaseOverflow(t *testing.T) {
|
|
messages := []message.ChatMessage{
|
|
{Role: "system", Content: strings.Repeat("系", 20)},
|
|
{Role: "user", Content: "最新问题"},
|
|
}
|
|
result := ApplyChatWindow(messages, 5)
|
|
if !result.BaseOverflow || !result.Overflow {
|
|
t.Fatalf("expected base overflow: %#v", result)
|
|
}
|
|
if len(result.ChatMessages) != len(messages) {
|
|
t.Fatalf("fixed/latest messages should remain: %#v", result.ChatMessages)
|
|
}
|
|
}
|
|
|
|
func TestApplyArkWindowKeepsToolContext(t *testing.T) {
|
|
messages := []*model.ChatCompletionMessage{
|
|
{Role: model.ChatMessageRoleSystem, Content: message.StringContent("system")},
|
|
{Role: model.ChatMessageRoleUser, Content: message.StringContent(strings.Repeat("旧", 40))},
|
|
{Role: model.ChatMessageRoleAssistant, Content: message.StringContent(strings.Repeat("旧", 40))},
|
|
{Role: model.ChatMessageRoleAssistant, ToolCalls: []*model.ToolCall{{ID: "call_1", Type: model.ToolTypeFunction, Function: model.FunctionCall{Name: "search", Arguments: `{"q":"x"}`}}}},
|
|
{Role: model.ChatMessageRoleTool, ToolCallID: "call_1", Content: message.StringContent("外部引用")},
|
|
}
|
|
result := ApplyArkWindow(messages, 35)
|
|
if result.Removed == 0 {
|
|
t.Fatalf("expected dialogue pruning: %#v", result)
|
|
}
|
|
var hasSystem, hasToolCall, hasTool bool
|
|
for _, msg := range result.ArkMessages {
|
|
if msg.Role == model.ChatMessageRoleSystem {
|
|
hasSystem = true
|
|
}
|
|
if len(msg.ToolCalls) > 0 {
|
|
hasToolCall = true
|
|
}
|
|
if msg.Role == model.ChatMessageRoleTool {
|
|
hasTool = true
|
|
}
|
|
}
|
|
if !hasSystem || !hasToolCall || !hasTool {
|
|
t.Fatalf("fixed partitions not preserved: %#v", result.ArkMessages)
|
|
}
|
|
}
|