+22
-11
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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 != "" {
|
||||
|
||||
@@ -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 == "" {
|
||||
|
||||
Reference in New Issue
Block a user