添加对话上下文窗口管理

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
2026-06-17 11:49:47 +08:00
co-authored by Claude
parent 54041e691b
commit 0b48dc0d7d
6 changed files with 365 additions and 15 deletions
+180
View File
@@ -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
}
+87
View File
@@ -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)
}
}