181 lines
4.4 KiB
Go
181 lines
4.4 KiB
Go
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
|
|
}
|