Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions backend/internal/application/conversation/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,7 @@ type Service struct {
snapshotCache sync.Map // conversationID (uint) → *cachedSnapshot
userMemCache sync.Map // userID (uint) → *cachedUserMemories
userSettingCache sync.Map // "userID:key" (string) → *cachedUserSetting
imageContextCache *preparedConversationImageCache
}

func (s *Service) llmAttribution() (string, string) {
Expand Down Expand Up @@ -133,6 +134,7 @@ type AttachmentInput struct {
RagOptOut bool // 用户是否关闭该文件的 RAG;RAG 段直接复用,无需重查 DB
ChunkCount int // 向量分块数;RAG 缓存 key 需要
Current bool // 是否为本轮用户显式上传的附件
MessageRole string
ContextMode string
}

Expand Down Expand Up @@ -253,6 +255,7 @@ func NewServiceWithRuntime(
storeProvider: appstorage.NewRuntimeProvider(cfg, nil),
logger: logger,
generationStreams: newGenerationStreamRegistry(cache, defaultGenerationStreamOptions()),
imageContextCache: defaultPreparedConversationImageCache(),
}
if extractSvc == nil {
extractSvc = extraction.NewServiceWithRuntime(cfg)
Expand Down
27 changes: 26 additions & 1 deletion backend/internal/application/conversation/service_branch.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@ import (
"go.uber.org/zap"
)

const estimatedConversationImageTokens int64 = 1024

type messageBranchState struct {
ExistingMessages []model.Message
ParentMessageID *uint
Expand Down Expand Up @@ -311,8 +313,9 @@ func truncateContextByTokenBudget(messages []model.Message, budgetTokens int, in
}
total := 0
cutFrom := len(messages)
imageTokenReserve := conversationImageTokenReserveByMessage(messages)
for i := len(messages) - 1; i >= 0; i-- {
msgTokens := int(estimateDomainMessageTokens(messages[i], includeReasoningContent))
msgTokens := int(estimateDomainMessageTokens(messages[i], includeReasoningContent) + imageTokenReserve[i])
if total+msgTokens > budgetTokens && cutFrom < len(messages) {
break
}
Expand All @@ -322,6 +325,28 @@ func truncateContextByTokenBudget(messages []model.Message, budgetTokens int, in
return messages[cutFrom:]
}

func conversationImageTokenReserveByMessage(messages []model.Message) map[int]int64 {
reserve := make(map[int]int64)
remaining := maxConversationImageContextCount
for messageIndex := len(messages) - 1; messageIndex >= 0 && remaining > 0; messageIndex-- {
message := messages[messageIndex]
if !strings.EqualFold(strings.TrimSpace(message.Role), "user") {
continue
}
refs := parseAttachmentSnapshotRefs(message.Attachments)
for attachmentIndex := len(refs) - 1; attachmentIndex >= 0 && remaining > 0; attachmentIndex-- {
ref := refs[attachmentIndex]
mimeType := firstNonEmptyString(ref.DetectedMIME, ref.MimeType)
if normalizeAttachmentKind(ref.Kind, mimeType) != "image" {
continue
}
reserve[messageIndex] += estimatedConversationImageTokens
remaining--
}
}
return reserve
}

func estimateDomainMessageTokens(message model.Message, includeReasoningContent bool) int64 {
tokens := estimateTokens(message.Content)
if includeReasoningContent && message.Role == "assistant" {
Expand Down
13 changes: 13 additions & 0 deletions backend/internal/application/conversation/service_branch_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -296,6 +296,19 @@ func TestTruncateContextByTokenBudgetCountsAssistantReasoningWhenEnabled(t *test
}
}

func TestTruncateContextByTokenBudgetReservesHistoricalImageTokens(t *testing.T) {
messages := []model.Message{
{ID: 1, Role: "user", Content: "first", Attachments: `[{"file_id":"image_1","kind":"image","mime_type":"image/png"}]`},
{ID: 2, Role: "assistant", Content: "ok"},
{ID: 3, Role: "user", Content: "next"},
}

got := truncateContextByTokenBudget(messages, 100, false)
if len(got) != 2 || got[0].ID != 2 || got[1].ID != 3 {
t.Fatalf("expected image token reserve to trim the oldest image turn, got %#v", got)
}
}

func TestBuildBranchMessagePathReusesExistingUserForAssistantRetry(t *testing.T) {
rootID := uint(1)
userID := uint(2)
Expand Down
63 changes: 54 additions & 9 deletions backend/internal/application/conversation/service_file_context.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,10 @@ const (
)

type attachmentSnapshotRef struct {
FileID string `json:"file_id"`
FileID string `json:"file_id"`
Kind string `json:"kind"`
MimeType string `json:"mime_type"`
DetectedMIME string `json:"detected_mime"`
}

type conversationFileContextPlan struct {
Expand Down Expand Up @@ -59,6 +62,17 @@ func collectConversationFileIDs(messages []model.Message, currentFileIDs []strin
}

func parseAttachmentSnapshotFileIDs(raw string) []string {
items := parseAttachmentSnapshotRefs(raw)
result := make([]string, 0, len(items))
for _, item := range items {
if fileID := strings.TrimSpace(item.FileID); fileID != "" {
result = append(result, fileID)
}
}
return result
}

func parseAttachmentSnapshotRefs(raw string) []attachmentSnapshotRef {
payload := strings.TrimSpace(raw)
if payload == "" || payload == "[]" {
return nil
Expand All @@ -67,13 +81,7 @@ func parseAttachmentSnapshotFileIDs(raw string) []string {
if err := json.Unmarshal([]byte(payload), &items); err != nil {
return nil
}
result := make([]string, 0, len(items))
for _, item := range items {
if fileID := strings.TrimSpace(item.FileID); fileID != "" {
result = append(result, fileID)
}
}
return result
return items
}

func filterCurrentAttachments(items []AttachmentInput) []AttachmentInput {
Expand All @@ -86,6 +94,29 @@ func filterCurrentAttachments(items []AttachmentInput) []AttachmentInput {
return result
}

func bindAttachmentMessageRoles(items []AttachmentInput, messages []model.Message) []AttachmentInput {
if len(items) == 0 || len(messages) == 0 {
return items
}
roles := make(map[string]string)
for _, message := range messages {
role := strings.ToLower(strings.TrimSpace(message.Role))
if role != "user" && role != "assistant" {
continue
}
for _, fileID := range parseAttachmentSnapshotFileIDs(message.Attachments) {
if role == "user" || roles[fileID] == "" {
roles[fileID] = role
}
}
}
result := append([]AttachmentInput(nil), items...)
for index := range result {
result[index].MessageRole = roles[strings.TrimSpace(result[index].FileID)]
}
return result
}

func filterAttachmentsByContextMode(items []AttachmentInput, contextMode string) []AttachmentInput {
result := make([]AttachmentInput, 0)
for _, item := range items {
Expand All @@ -105,6 +136,9 @@ func isStableTextAttachment(item AttachmentInput) bool {

func shouldShowAttachmentProcessTrace(items []AttachmentInput) bool {
for _, item := range items {
if strings.EqualFold(strings.TrimSpace(item.ContextMode), fileContextModeDirectImage) && !item.Current {
continue
}
if item.Current {
return true
}
Expand All @@ -115,6 +149,17 @@ func shouldShowAttachmentProcessTrace(items []AttachmentInput) bool {
return false
}

func attachmentProcessTraceItems(items []AttachmentInput) []AttachmentInput {
result := make([]AttachmentInput, 0, len(items))
for _, item := range items {
if strings.EqualFold(strings.TrimSpace(item.ContextMode), fileContextModeDirectImage) && !item.Current {
continue
}
result = append(result, item)
}
return result
}

func buildConversationFileContextPlan(
attachments []AttachmentInput,
fileMode string,
Expand All @@ -128,7 +173,7 @@ func buildConversationFileContextPlan(
}
for _, item := range attachments {
kind := normalizeAttachmentKind(item.Kind, item.DetectedMIME)
if kind == "image" && item.Current {
if kind == "image" && (item.Current || strings.EqualFold(strings.TrimSpace(item.MessageRole), "user")) {
item.ContextMode = fileContextModeDirectImage
plan.Attachments = append(plan.Attachments, item)
plan.FullAttachments = append(plan.FullAttachments, item)
Expand Down
Loading
Loading