From 0b48dc0d7d5053cdd17bcb233092918505b5ce5d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=97=A0=E9=97=BB=E9=A3=8E?= Date: Wed, 17 Jun 2026 11:49:47 +0800 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0=E5=AF=B9=E8=AF=9D=E4=B8=8A?= =?UTF-8?q?=E4=B8=8B=E6=96=87=E7=AA=97=E5=8F=A3=E7=AE=A1=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude --- config/config.go | 33 ++++--- contextwindow/window.go | 180 +++++++++++++++++++++++++++++++++++ contextwindow/window_test.go | 87 +++++++++++++++++ main_test.go | 15 +++ server/handlers.go | 21 +++- stream/tokens.go | 44 +++++++++ 6 files changed, 365 insertions(+), 15 deletions(-) create mode 100644 contextwindow/window.go create mode 100644 contextwindow/window_test.go diff --git a/config/config.go b/config/config.go index 6354834..9356a35 100644 --- a/config/config.go +++ b/config/config.go @@ -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 diff --git a/contextwindow/window.go b/contextwindow/window.go new file mode 100644 index 0000000..d66c5a9 --- /dev/null +++ b/contextwindow/window.go @@ -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 +} diff --git a/contextwindow/window_test.go b/contextwindow/window_test.go new file mode 100644 index 0000000..9bdbab9 --- /dev/null +++ b/contextwindow/window_test.go @@ -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) + } +} diff --git a/main_test.go b/main_test.go index 695edd0..9cb7879 100644 --- a/main_test.go +++ b/main_test.go @@ -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, diff --git a/server/handlers.go b/server/handlers.go index 04afe55..1cbe74f 100644 --- a/server/handlers.go +++ b/server/handlers.go @@ -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 != "" { diff --git a/stream/tokens.go b/stream/tokens.go index 1f62461..7750185 100644 --- a/stream/tokens.go +++ b/stream/tokens.go @@ -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 == "" {