@@ -0,0 +1,180 @@
|
||||
package contextwindow
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"aichat/message"
|
||||
"aichat/stream"
|
||||
|
||||
"github.com/volcengine/volcengine-go-sdk/service/arkruntime/model"
|
||||
)
|
||||
|
||||
type Result struct {
|
||||
MaxTokens int
|
||||
BeforeTokens int
|
||||
AfterTokens int
|
||||
Removed int
|
||||
Overflow bool
|
||||
BaseOverflow bool
|
||||
ChatMessages []message.ChatMessage
|
||||
ArkMessages []*model.ChatCompletionMessage
|
||||
}
|
||||
|
||||
type chatItem struct {
|
||||
msg message.ChatMessage
|
||||
tokens int
|
||||
fixed bool
|
||||
protected bool
|
||||
}
|
||||
|
||||
type arkItem struct {
|
||||
msg *model.ChatCompletionMessage
|
||||
tokens int
|
||||
fixed bool
|
||||
}
|
||||
|
||||
func ApplyChatWindow(messages []message.ChatMessage, maxTokens int) Result {
|
||||
before := stream.EstimateChatMessagesTokens(messages)
|
||||
result := Result{MaxTokens: maxTokens, BeforeTokens: before, AfterTokens: before, ChatMessages: append([]message.ChatMessage(nil), messages...)}
|
||||
if maxTokens <= 0 || before <= maxTokens {
|
||||
result.Overflow = maxTokens > 0 && before > maxTokens
|
||||
return result
|
||||
}
|
||||
|
||||
items := make([]chatItem, 0, len(messages))
|
||||
lastUser := -1
|
||||
baseTokens := 0
|
||||
for i, msg := range messages {
|
||||
item := chatItem{msg: msg, tokens: stream.EstimateChatMessagesTokens([]message.ChatMessage{msg}), fixed: isFixedChatMessage(msg)}
|
||||
if item.fixed {
|
||||
baseTokens += item.tokens
|
||||
} else if strings.EqualFold(msg.Role, "user") {
|
||||
lastUser = i
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
if lastUser >= 0 {
|
||||
items[lastUser].protected = true
|
||||
}
|
||||
if baseTokens > maxTokens {
|
||||
result.BaseOverflow = true
|
||||
}
|
||||
|
||||
result.ChatMessages, result.Removed = pruneChatItems(items, maxTokens)
|
||||
result.AfterTokens = stream.EstimateChatMessagesTokens(result.ChatMessages)
|
||||
result.Overflow = result.AfterTokens > maxTokens
|
||||
return result
|
||||
}
|
||||
|
||||
func ApplyArkWindow(messages []*model.ChatCompletionMessage, maxTokens int) Result {
|
||||
before := stream.EstimateArkMessagesTokens(messages)
|
||||
result := Result{MaxTokens: maxTokens, BeforeTokens: before, AfterTokens: before, ArkMessages: append([]*model.ChatCompletionMessage(nil), messages...)}
|
||||
if maxTokens <= 0 || before <= maxTokens {
|
||||
result.Overflow = maxTokens > 0 && before > maxTokens
|
||||
return result
|
||||
}
|
||||
|
||||
items := make([]arkItem, 0, len(messages))
|
||||
lastUser := -1
|
||||
baseTokens := 0
|
||||
for i, msg := range messages {
|
||||
item := arkItem{msg: msg, tokens: stream.EstimateArkMessagesTokens([]*model.ChatCompletionMessage{msg}), fixed: isFixedArkMessage(msg)}
|
||||
if item.fixed {
|
||||
baseTokens += item.tokens
|
||||
} else if msg != nil && msg.Role == model.ChatMessageRoleUser {
|
||||
lastUser = i
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
if lastUser >= 0 {
|
||||
items[lastUser].fixed = true
|
||||
}
|
||||
if baseTokens > maxTokens {
|
||||
result.BaseOverflow = true
|
||||
}
|
||||
|
||||
result.ArkMessages, result.Removed = pruneArkItems(items, maxTokens)
|
||||
result.AfterTokens = stream.EstimateArkMessagesTokens(result.ArkMessages)
|
||||
result.Overflow = result.AfterTokens > maxTokens
|
||||
return result
|
||||
}
|
||||
|
||||
func isFixedChatMessage(msg message.ChatMessage) bool {
|
||||
return msg.Hidden || strings.EqualFold(msg.Role, "system")
|
||||
}
|
||||
|
||||
func isFixedArkMessage(msg *model.ChatCompletionMessage) bool {
|
||||
if msg == nil {
|
||||
return false
|
||||
}
|
||||
if msg.Role == model.ChatMessageRoleSystem || msg.Role == model.ChatMessageRoleTool {
|
||||
return true
|
||||
}
|
||||
return len(msg.ToolCalls) > 0
|
||||
}
|
||||
|
||||
func pruneChatItems(items []chatItem, maxTokens int) ([]message.ChatMessage, int) {
|
||||
removed := make([]bool, len(items))
|
||||
total := 0
|
||||
for _, item := range items {
|
||||
total += item.tokens
|
||||
}
|
||||
removedCount := 0
|
||||
for total > maxTokens {
|
||||
idx := -1
|
||||
for i, item := range items {
|
||||
if removed[i] || item.fixed || item.protected {
|
||||
continue
|
||||
}
|
||||
idx = i
|
||||
break
|
||||
}
|
||||
if idx == -1 {
|
||||
break
|
||||
}
|
||||
removed[idx] = true
|
||||
total -= items[idx].tokens
|
||||
removedCount++
|
||||
}
|
||||
|
||||
messages := make([]message.ChatMessage, 0, len(items)-removedCount)
|
||||
for i, item := range items {
|
||||
if !removed[i] {
|
||||
messages = append(messages, item.msg)
|
||||
}
|
||||
}
|
||||
return messages, removedCount
|
||||
}
|
||||
|
||||
func pruneArkItems(items []arkItem, maxTokens int) ([]*model.ChatCompletionMessage, int) {
|
||||
removed := make([]bool, len(items))
|
||||
total := 0
|
||||
for _, item := range items {
|
||||
total += item.tokens
|
||||
}
|
||||
removedCount := 0
|
||||
for total > maxTokens {
|
||||
idx := -1
|
||||
for i, item := range items {
|
||||
if removed[i] || item.fixed {
|
||||
continue
|
||||
}
|
||||
idx = i
|
||||
break
|
||||
}
|
||||
if idx == -1 {
|
||||
break
|
||||
}
|
||||
removed[idx] = true
|
||||
total -= items[idx].tokens
|
||||
removedCount++
|
||||
}
|
||||
|
||||
messages := make([]*model.ChatCompletionMessage, 0, len(items)-removedCount)
|
||||
for i, item := range items {
|
||||
if !removed[i] {
|
||||
messages = append(messages, item.msg)
|
||||
}
|
||||
}
|
||||
return messages, removedCount
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user