添加对话上下文窗口管理

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
+22 -11
View File
@@ -13,6 +13,7 @@ import (
const (
defaultOpenAIBaseURL = "https://ark.cn-beijing.volces.com/api/v3"
defaultOpenAITimeout = 120
defaultContextWindowTokens = 262144
defaultToolRouterTimeout = 30
defaultToolRouterMaxTokens = 512
defaultToolRouterSystemText = `你可以按需直接调用可用工具来回答用户问题。
@@ -24,13 +25,14 @@ const (
)
type OpenAIConfig struct {
Name string `yaml:"name" json:"name"`
Active bool `yaml:"active,omitempty" json:"active"`
APIKey string `yaml:"api_key" json:"-"`
BaseURL string `yaml:"base_url" json:"base_url"`
Model string `yaml:"model" json:"model"`
Timeout int `yaml:"timeout" json:"timeout"`
ParseThinkTags *bool `yaml:"parse_think_tags,omitempty" json:"parse_think_tags,omitempty"`
Name string `yaml:"name" json:"name"`
Active bool `yaml:"active,omitempty" json:"active"`
APIKey string `yaml:"api_key" json:"-"`
BaseURL string `yaml:"base_url" json:"base_url"`
Model string `yaml:"model" json:"model"`
Timeout int `yaml:"timeout" json:"timeout"`
ContextWindowTokens int `yaml:"context_window_tokens" json:"context_window_tokens"`
ParseThinkTags *bool `yaml:"parse_think_tags,omitempty" json:"parse_think_tags,omitempty"`
}
type OpenAIConfigs []OpenAIConfig
@@ -87,10 +89,11 @@ type Config struct {
func defaultOpenAIConfig() OpenAIConfig {
return OpenAIConfig{
Name: "default",
Active: true,
BaseURL: defaultOpenAIBaseURL,
Timeout: defaultOpenAITimeout,
Name: "default",
Active: true,
BaseURL: defaultOpenAIBaseURL,
Timeout: defaultOpenAITimeout,
ContextWindowTokens: defaultContextWindowTokens,
}
}
@@ -209,6 +212,10 @@ func ensureFile(path string) error {
return Write(path, cfg)
}
func NormalizeOpenAIConfigs(cfg *Config) (bool, error) {
return normalizeOpenAIConfigs(cfg)
}
func normalizeOpenAIConfigs(cfg *Config) (bool, error) {
changed := false
if len(cfg.OpenAI) == 0 {
@@ -245,6 +252,10 @@ func normalizeOpenAIConfigs(cfg *Config) (bool, error) {
profile.Timeout = defaultOpenAITimeout
changed = true
}
if profile.ContextWindowTokens <= 0 {
profile.ContextWindowTokens = defaultContextWindowTokens
changed = true
}
if profile.Active {
if activeIndex == -1 {
activeIndex = i
+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)
}
}
+15
View File
@@ -51,6 +51,21 @@ func TestNormalizeToolRouterConfigDefaults(t *testing.T) {
}
}
func TestNormalizeOpenAIConfigDefaultsContextWindow(t *testing.T) {
cfg := config.Default()
cfg.OpenAI[0].ContextWindowTokens = 0
changed, err := config.NormalizeOpenAIConfigs(&cfg)
if err != nil {
t.Fatal(err)
}
if !changed {
t.Fatal("expected context window default to change config")
}
if cfg.OpenAI[0].ContextWindowTokens != 262144 {
t.Fatalf("context_window_tokens = %d", cfg.OpenAI[0].ContextWindowTokens)
}
}
func TestNormalizeToolRouterConfigAddsTimeBeforeSQL(t *testing.T) {
cfg := &config.Config{ToolRouter: config.ToolRouterConfig{
Enabled: true,
+17 -4
View File
@@ -10,6 +10,7 @@ import (
"strings"
"time"
"aichat/contextwindow"
"aichat/conversation"
"aichat/llm"
"aichat/message"
@@ -162,19 +163,31 @@ func (s *Server) chatHandler(c *gin.Context) {
usage := stream.NewTracker()
ctx = stream.ContextWithTracker(ctx, usage)
contextMessages := req.Messages
chatWindow := contextwindow.ApplyChatWindow(req.Messages, profile.Config.ContextWindowTokens)
if chatWindow.Removed > 0 || chatWindow.Overflow {
emitTrace("context_window", "chat", "success", "已清理对话历史上下文", map[string]any{"max_tokens": chatWindow.MaxTokens, "before_tokens": chatWindow.BeforeTokens, "after_tokens": chatWindow.AfterTokens, "removed_messages": chatWindow.Removed, "overflow": chatWindow.Overflow, "base_overflow": chatWindow.BaseOverflow})
}
contextMessages = chatWindow.ChatMessages
// 用 Function Calling 工具循环替代旧的路由+隐藏上下文机制
messages, err := toolrouter.RunAgentToolLoop(ctx, s.toolRouterState, profile, req.Messages, s.searchState, s.sqlState, emit)
messages, err := toolrouter.RunAgentToolLoop(ctx, s.toolRouterState, profile, contextMessages, s.searchState, s.sqlState, emit)
if err != nil {
fmt.Fprintln(os.Stderr, "Agent 工具循环失败:", err)
messages, err = message.BuildArkMessages(req.Messages)
messages, err = message.BuildArkMessages(contextMessages)
if err != nil {
emitError(err)
return
}
}
promptTokens := stream.EstimateChatMessagesTokens(req.Messages)
arkWindow := contextwindow.ApplyArkWindow(messages, profile.Config.ContextWindowTokens)
if arkWindow.Removed > 0 || arkWindow.Overflow {
emitTrace("context_window", "model", "success", "已清理最终模型上下文", map[string]any{"max_tokens": arkWindow.MaxTokens, "before_tokens": arkWindow.BeforeTokens, "after_tokens": arkWindow.AfterTokens, "removed_messages": arkWindow.Removed, "overflow": arkWindow.Overflow, "base_overflow": arkWindow.BaseOverflow})
}
messages = arkWindow.ArkMessages
promptTokens := arkWindow.AfterTokens
if llm.IsOllamaProfile(profile) && message.HasImageMessage(req.Messages) {
if llm.IsOllamaProfile(profile) && message.HasImageMessage(contextMessages) {
emitTrace("model", "request", "running", "正在通过 Ollama 原生接口调用视觉模型", nil)
err = stream.StreamOllamaChat(ctx, profile, messages, promptTokens, usage, emit, func(content string) {
if req.ConversationID != "" {
+44
View File
@@ -8,6 +8,8 @@ import (
"unicode"
"aichat/message"
"github.com/volcengine/volcengine-go-sdk/service/arkruntime/model"
)
type Stats struct {
@@ -97,6 +99,48 @@ func EstimateChatMessagesTokens(messages []message.ChatMessage) int {
return total
}
func EstimateArkMessagesTokens(messages []*model.ChatCompletionMessage) int {
total := 0
for _, msg := range messages {
if msg == nil {
continue
}
total += EstimateTokenCount(string(msg.Role)) + estimateArkContentTokens(msg.Content) + 4
if msg.ToolCallID != "" {
total += EstimateTokenCount(msg.ToolCallID) + 2
}
for _, call := range msg.ToolCalls {
if call == nil {
continue
}
total += EstimateTokenCount(call.ID) + EstimateTokenCount(call.Function.Name) + EstimateTokenCount(call.Function.Arguments) + 8
}
}
return total
}
func estimateArkContentTokens(content *model.ChatCompletionMessageContent) int {
if content == nil {
return 0
}
if content.StringValue != nil {
return EstimateTokenCount(*content.StringValue)
}
total := 0
for _, part := range content.ListValue {
if part == nil {
continue
}
if part.Text != "" {
total += EstimateTokenCount(part.Text)
}
if part.ImageURL != nil {
total += 85
}
}
return total
}
func EstimateTokenCount(text string) int {
text = strings.TrimSpace(text)
if text == "" {