diff --git a/conversation/store.go b/conversation/store.go index a5162a4..874265d 100644 --- a/conversation/store.go +++ b/conversation/store.go @@ -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 { diff --git a/conversation/store_test.go b/conversation/store_test.go new file mode 100644 index 0000000..23263d5 --- /dev/null +++ b/conversation/store_test.go @@ -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]) + } +} diff --git a/message/types.go b/message/types.go index 448bf38..bd50b8c 100644 --- a/message/types.go +++ b/message/types.go @@ -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"` } diff --git a/server/handlers.go b/server/handlers.go index 1cbe74f..9be2bfc 100644 --- a/server/handlers.go +++ b/server/handlers.go @@ -26,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{ @@ -89,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 @@ -118,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 @@ -163,8 +211,8 @@ 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) + 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}) } diff --git a/templates/chat.html b/templates/chat.html index e303392..8cf0f51 100644 --- a/templates/chat.html +++ b/templates/chat.html @@ -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();