Compare commits

...
2 Commits
Author SHA1 Message Date
kevinandClaude 2c4d4af070 添加会话级固定预设
Co-Authored-By: Claude <noreply@anthropic.com>
2026-06-17 12:06:13 +08:00
kevinandClaude 0b48dc0d7d 添加对话上下文窗口管理
Co-Authored-By: Claude <noreply@anthropic.com>
2026-06-17 11:49:47 +08:00
10 changed files with 517 additions and 28 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)
}
}
+21 -5
View File
@@ -30,11 +30,16 @@ func (s *Store) path(id string) string {
}
func (s *Store) Create() (*message.Conversation, error) {
return s.CreateWithPreset("")
}
func (s *Store) CreateWithPreset(preset string) (*message.Conversation, error) {
conv := &message.Conversation{
ID: utils.NewUUID(),
Title: "新对话",
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
ID: utils.NewUUID(),
Title: "新对话",
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
PresetPrompt: strings.TrimSpace(preset),
}
if err := s.Save(conv); err != nil {
return nil, err
@@ -124,7 +129,7 @@ func SaveMessages(store *Store, id string, messages []message.ChatMessage, assis
if err != nil {
return err
}
conv.Messages = append([]message.ChatMessage(nil), messages...)
conv.Messages = filterPresetMessages(messages)
conv.Messages = append(conv.Messages, message.ChatMessage{Role: "assistant", Content: assistantContent})
if conv.Title == "" || conv.Title == "新对话" {
conv.Title = GenTitle(conv.Messages)
@@ -132,6 +137,17 @@ func SaveMessages(store *Store, id string, messages []message.ChatMessage, assis
return store.Save(conv)
}
func filterPresetMessages(messages []message.ChatMessage) []message.ChatMessage {
filtered := make([]message.ChatMessage, 0, len(messages))
for _, msg := range messages {
if msg.Hidden && strings.EqualFold(msg.Role, "system") {
continue
}
filtered = append(filtered, msg)
}
return filtered
}
func GenTitle(messages []message.ChatMessage) string {
for _, m := range messages {
if m.Hidden {
+70
View File
@@ -0,0 +1,70 @@
package conversation
import (
"testing"
"aichat/message"
)
func TestCreateWithPresetSavesPresetPrompt(t *testing.T) {
store := NewStore(t.TempDir())
conv, err := store.CreateWithPreset(" 你是测试助手 ")
if err != nil {
t.Fatal(err)
}
if conv.PresetPrompt != "你是测试助手" {
t.Fatalf("preset_prompt = %q", conv.PresetPrompt)
}
loaded, err := store.Get(conv.ID)
if err != nil {
t.Fatal(err)
}
if loaded.PresetPrompt != "你是测试助手" {
t.Fatalf("loaded preset_prompt = %q", loaded.PresetPrompt)
}
if len(loaded.Messages) != 0 {
t.Fatalf("preset should not be stored in messages: %#v", loaded.Messages)
}
}
func TestCreateKeepsEmptyPresetPrompt(t *testing.T) {
store := NewStore(t.TempDir())
conv, err := store.Create()
if err != nil {
t.Fatal(err)
}
if conv.PresetPrompt != "" {
t.Fatalf("preset_prompt = %q", conv.PresetPrompt)
}
}
func TestSaveMessagesKeepsPresetPromptAndFiltersPresetMessages(t *testing.T) {
store := NewStore(t.TempDir())
conv, err := store.CreateWithPreset("你是测试助手")
if err != nil {
t.Fatal(err)
}
err = SaveMessages(store, conv.ID, []message.ChatMessage{
{Role: "system", Content: "你是测试助手", Hidden: true},
{Role: "user", Content: "你好"},
}, "你好,我在")
if err != nil {
t.Fatal(err)
}
loaded, err := store.Get(conv.ID)
if err != nil {
t.Fatal(err)
}
if loaded.PresetPrompt != "你是测试助手" {
t.Fatalf("preset_prompt = %q", loaded.PresetPrompt)
}
if len(loaded.Messages) != 2 {
t.Fatalf("messages length = %d: %#v", len(loaded.Messages), loaded.Messages)
}
if loaded.Messages[0].Role != "user" || loaded.Messages[0].Content != "你好" {
t.Fatalf("unexpected first message: %#v", loaded.Messages[0])
}
if loaded.Messages[1].Role != "assistant" || loaded.Messages[1].Content != "你好,我在" {
t.Fatalf("unexpected assistant message: %#v", loaded.Messages[1])
}
}
+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,
+6 -5
View File
@@ -18,9 +18,10 @@ type ChatRequest struct {
}
type Conversation struct {
ID string `json:"id"`
Title string `json:"title"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
Messages []ChatMessage `json:"messages,omitempty"`
ID string `json:"id"`
Title string `json:"title"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
PresetPrompt string `json:"preset_prompt,omitempty"`
Messages []ChatMessage `json:"messages,omitempty"`
}
+66 -5
View File
@@ -10,6 +10,7 @@ import (
"strings"
"time"
"aichat/contextwindow"
"aichat/conversation"
"aichat/llm"
"aichat/message"
@@ -25,6 +26,11 @@ type activeProfileRequest struct {
Name string `json:"name"`
}
type createConversationRequest struct {
Preset string `json:"preset"`
PresetPrompt string `json:"preset_prompt"`
}
func (s *Server) indexHandler(c *gin.Context) {
profile := s.aiState.ActiveProfile()
c.HTML(http.StatusOK, "chat.html", gin.H{
@@ -88,7 +94,15 @@ func (s *Server) listConversationsHandler(c *gin.Context) {
}
func (s *Server) createConversationHandler(c *gin.Context) {
conv, err := s.store.Create()
var req createConversationRequest
if c.Request.Body != nil {
_ = c.ShouldBindJSON(&req)
}
preset := req.PresetPrompt
if strings.TrimSpace(preset) == "" {
preset = req.Preset
}
conv, err := s.store.CreateWithPreset(preset)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "创建对话失败: " + err.Error()})
return
@@ -117,6 +131,41 @@ func (s *Server) deleteConversationHandler(c *gin.Context) {
c.Status(http.StatusNoContent)
}
func conversationPresetPrompt(store *conversation.Store, id string) string {
id = strings.TrimSpace(id)
if id == "" || store == nil {
return ""
}
conv, err := store.Get(id)
if err != nil {
return ""
}
return strings.TrimSpace(conv.PresetPrompt)
}
func conversationContextMessages(messages []message.ChatMessage, preset string) []message.ChatMessage {
preset = strings.TrimSpace(preset)
if preset == "" {
return append([]message.ChatMessage(nil), messages...)
}
cleaned := filterRequestPresetMessages(messages)
contextMessages := make([]message.ChatMessage, 0, len(cleaned)+1)
contextMessages = append(contextMessages, message.ChatMessage{Role: "system", Content: preset, Hidden: true})
contextMessages = append(contextMessages, cleaned...)
return contextMessages
}
func filterRequestPresetMessages(messages []message.ChatMessage) []message.ChatMessage {
filtered := make([]message.ChatMessage, 0, len(messages))
for _, msg := range messages {
if msg.Hidden && strings.EqualFold(msg.Role, "system") {
continue
}
filtered = append(filtered, msg)
}
return filtered
}
// chatHandler 流式 SSE 对话接口
func (s *Server) chatHandler(c *gin.Context) {
var req message.ChatRequest
@@ -162,19 +211,31 @@ func (s *Server) chatHandler(c *gin.Context) {
usage := stream.NewTracker()
ctx = stream.ContextWithTracker(ctx, usage)
contextMessages := conversationContextMessages(req.Messages, conversationPresetPrompt(s.store, req.ConversationID))
chatWindow := contextwindow.ApplyChatWindow(contextMessages, 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 == "" {
+6 -2
View File
@@ -771,7 +771,12 @@ async function loadConversationList() {
}
async function createConversation() {
const res = await fetch('/api/conversations', { method: 'POST' });
const preset = getPresetPrompt();
const res = await fetch('/api/conversations', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ preset_prompt: preset }),
});
if (!res.ok) {
const err = await res.json().catch(() => ({ error: '创建对话失败' }));
throw new Error(err.error || '创建对话失败');
@@ -1075,7 +1080,6 @@ async function sendMessage() {
try {
if (!currentConvId) {
ensurePresetLoaded();
const conv = await createConversation();
currentConvId = conv.id;
await loadConversationList();