From f1f3a0048c5b58ef9da1b6ca1ad7970b9e23f0c6 Mon Sep 17 00:00:00 2001 From: Omar Alani Date: Mon, 3 Aug 2026 22:04:21 -0500 Subject: [PATCH 1/6] feat(terminal): support image prompt attachments --- cmd/librecode/cli_helpers_internal_test.go | 2 +- cmd/librecode/prompt.go | 1 + internal/agenttask/runtime_runner.go | 2 +- internal/assistant/context_compaction_test.go | 2 +- internal/assistant/lifecycle.go | 7 + .../lifecyclepayload/behavior_test.go | 2 +- .../lifecyclepayload/lifecyclepayload.go | 15 +- .../lifecyclepayload/lifecyclepayload_test.go | 3 +- internal/assistant/llm_conversion.go | 35 ++- .../assistant/llm_conversion_internal_test.go | 34 +++ internal/assistant/prompt_images.go | 231 ++++++++++++++++ .../assistant/prompt_images_internal_test.go | 112 ++++++++ internal/assistant/runtime.go | 34 ++- .../runtime_context_internal_test.go | 2 +- internal/assistant/runtime_entries.go | 33 ++- internal/assistant/runtime_model.go | 52 +++- internal/assistant/runtime_persist.go | 7 +- internal/assistant/runtime_test.go | 131 ++++++++- .../test_message_helpers_internal_test.go | 2 +- internal/compaction/plan.go | 48 +++- internal/compaction/plan_internal_test.go | 50 +++- internal/contextwindow/tokens.go | 53 +++- internal/contextwindow/usage_internal_test.go | 20 +- .../contextwindow/usage_led_internal_test.go | 2 +- internal/database/entity.go | 61 +++-- .../00013_add_session_message_parts.sql | 28 ++ internal/database/migrations_test.go | 43 +++ .../repository_helpers_internal_test.go | 2 +- .../session_compaction_input_internal_test.go | 2 +- internal/database/session_entry_repository.go | 2 +- .../database/session_message_parts_test.go | 251 ++++++++++++++++++ .../database/session_message_repository.go | 210 +++++++++++++-- internal/database/session_repository_test.go | 18 +- internal/database/session_store.go | 29 +- internal/database/session_usage_test.go | 2 +- .../sqlite_contention_internal_test.go | 8 +- .../database/task_validation_internal_test.go | 4 +- internal/database/test_helpers_test.go | 2 +- internal/database/validation.go | 109 ++++++++ internal/extension/lua_values.go | 7 + internal/model/message_filter.go | 2 +- internal/model/messages.go | 35 ++- internal/model/messages_test.go | 23 +- internal/provider/anthropic.go | 27 +- .../anthropic_mapping_internal_test.go | 50 +++- internal/provider/image_content.go | 159 +++++++++++ .../provider/image_content_internal_test.go | 194 ++++++++++++++ internal/provider/messages.go | 22 +- internal/provider/messages_internal_test.go | 52 +++- internal/provider/openai_chat.go | 27 +- .../openai_chat_payload_internal_test.go | 54 +++- internal/provider/openai_responses.go | 25 +- internal/terminal/agent_tasks.go | 17 +- .../agent_tasks_behavior_internal_test.go | 18 +- .../agent_tasks_live_internal_test.go | 2 +- internal/terminal/app.go | 106 +++++--- .../terminal/async_events_internal_test.go | 11 +- internal/terminal/attachment_actions.go | 60 +++++ internal/terminal/attachments.go | 249 +++++++++++++++++ .../terminal/attachments_internal_test.go | 241 +++++++++++++++++ .../terminal/auth_commands_internal_test.go | 2 +- internal/terminal/clipboard.go | 20 ++ internal/terminal/clipboard_internal_test.go | 1 + .../compact_commands_internal_test.go | 8 +- .../extension_events_internal_test.go | 44 +++ internal/terminal/input.go | 47 +++- internal/terminal/input_escape.go | 3 +- internal/terminal/interrupt_internal_test.go | 1 + internal/terminal/keybindings.go | 6 + internal/terminal/message_render.go | 26 +- .../model_test_helpers_internal_test.go | 18 ++ .../panel_session_selection_internal_test.go | 2 +- .../panel_test_helpers_internal_test.go | 2 +- internal/terminal/panel_tree_internal_test.go | 4 +- .../terminal/prompt_cancel_internal_test.go | 4 +- internal/terminal/prompt_history.go | 54 +++- .../terminal/prompt_history_internal_test.go | 30 +++ internal/terminal/prompt_queue.go | 61 +++-- .../terminal/prompt_queue_internal_test.go | 95 ++++++- .../terminal/prompt_response_internal_test.go | 2 +- internal/terminal/prompt_send.go | 18 +- .../terminal/prompt_send_internal_test.go | 34 ++- internal/terminal/prompt_submit.go | 63 +++-- internal/terminal/render_composer.go | 102 ++++++- internal/terminal/render_internal_test.go | 33 ++- .../terminal/render_parity_internal_test.go | 2 +- .../terminal/running_tools_internal_test.go | 9 +- .../session_commands_internal_test.go | 2 +- internal/terminal/session_view.go | 119 +++++---- .../workflow_summary_internal_test.go | 4 +- 90 files changed, 3498 insertions(+), 350 deletions(-) create mode 100644 internal/assistant/prompt_images.go create mode 100644 internal/assistant/prompt_images_internal_test.go create mode 100644 internal/database/migrations/00013_add_session_message_parts.sql create mode 100644 internal/database/session_message_parts_test.go create mode 100644 internal/provider/image_content.go create mode 100644 internal/provider/image_content_internal_test.go create mode 100644 internal/terminal/attachment_actions.go create mode 100644 internal/terminal/attachments.go create mode 100644 internal/terminal/attachments_internal_test.go diff --git a/cmd/librecode/cli_helpers_internal_test.go b/cmd/librecode/cli_helpers_internal_test.go index 5de33cc9..44f58404 100644 --- a/cmd/librecode/cli_helpers_internal_test.go +++ b/cmd/librecode/cli_helpers_internal_test.go @@ -111,7 +111,7 @@ func TestPrintSessionSummaryAndEntry(t *testing.T) { Role: database.RoleUser, Content: "message text", Provider: "", - Model: "", + Model: "", Parts: nil, }, Summary: "", ToolStatus: "", diff --git a/cmd/librecode/prompt.go b/cmd/librecode/prompt.go index 4ab252c8..0b18bd85 100644 --- a/cmd/librecode/prompt.go +++ b/cmd/librecode/prompt.go @@ -199,6 +199,7 @@ func buildPromptRequest(cwd, message string, options promptRunOptions) *assistan ParentEntryID: nil, SessionID: options.SessionID, CWD: cwd, + Images: nil, Text: message, Name: options.SessionName, ResumeLatest: options.Resume, diff --git a/internal/agenttask/runtime_runner.go b/internal/agenttask/runtime_runner.go index 75a7c11d..d4656ed6 100644 --- a/internal/agenttask/runtime_runner.go +++ b/internal/agenttask/runtime_runner.go @@ -72,7 +72,7 @@ func (runner *RuntimeRunner) Run( }, OnRetry: nil, OnUserEntry: nil, ParentEntryID: nil, SessionID: task.ChildSessionID, CWD: session.CWD, Text: task.Prompt, - Name: "", ResumeLatest: false, HideUserPrompt: false, + Images: nil, Name: "", ResumeLatest: false, HideUserPrompt: false, }) usageJSON, usageErr := agentUsageJSON(response, metrics.Snapshot()) diff --git a/internal/assistant/context_compaction_test.go b/internal/assistant/context_compaction_test.go index 84ee76bd..718c36e7 100644 --- a/internal/assistant/context_compaction_test.go +++ b/internal/assistant/context_compaction_test.go @@ -478,7 +478,7 @@ func appendRuntimeTestMessage( Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) diff --git a/internal/assistant/lifecycle.go b/internal/assistant/lifecycle.go index 5e5deeaf..b427d7fa 100644 --- a/internal/assistant/lifecycle.go +++ b/internal/assistant/lifecycle.go @@ -273,6 +273,7 @@ func lifecyclePromptRequest(request *PromptRequest) *lifecyclepayload.PromptRequ if request == nil { return &lifecyclepayload.PromptRequest{ ParentEntryID: nil, + Attachments: nil, CWD: "", Name: "", SessionID: "", @@ -281,8 +282,14 @@ func lifecyclePromptRequest(request *PromptRequest) *lifecyclepayload.PromptRequ } } + attachments := make([]map[string]any, len(request.Images)) + for index := range request.Images { + attachments[index] = imageMetadata(request.Images[index]) + } + return &lifecyclepayload.PromptRequest{ ParentEntryID: request.ParentEntryID, + Attachments: attachments, CWD: request.CWD, Name: request.Name, SessionID: request.SessionID, diff --git a/internal/assistant/lifecyclepayload/behavior_test.go b/internal/assistant/lifecyclepayload/behavior_test.go index c5381ac3..404dbcf9 100644 --- a/internal/assistant/lifecyclepayload/behavior_test.go +++ b/internal/assistant/lifecyclepayload/behavior_test.go @@ -202,7 +202,7 @@ func messageEntity(role database.Role, content string) database.MessageEntity { Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, } } diff --git a/internal/assistant/lifecyclepayload/lifecyclepayload.go b/internal/assistant/lifecyclepayload/lifecyclepayload.go index ccfca71c..82417977 100644 --- a/internal/assistant/lifecyclepayload/lifecyclepayload.go +++ b/internal/assistant/lifecyclepayload/lifecyclepayload.go @@ -57,6 +57,7 @@ type PromptRequest struct { Name string SessionID string Text string + Attachments []map[string]any ResumeLatest bool } @@ -145,12 +146,14 @@ func Prompt(request *PromptRequest) map[string]any { } return map[string]any{ - CWDKey: request.CWD, - ToolNameKey: request.Name, - ParentEntryIDKey: stringPtrValue(request.ParentEntryID), - PromptKey: request.Text, - "resume_latest": request.ResumeLatest, - SessionIDKey: request.SessionID, + "attachments": request.Attachments, + "attachment_count": len(request.Attachments), + CWDKey: request.CWD, + ToolNameKey: request.Name, + ParentEntryIDKey: stringPtrValue(request.ParentEntryID), + PromptKey: request.Text, + "resume_latest": request.ResumeLatest, + SessionIDKey: request.SessionID, } } diff --git a/internal/assistant/lifecyclepayload/lifecyclepayload_test.go b/internal/assistant/lifecyclepayload/lifecyclepayload_test.go index 5112e008..1c145636 100644 --- a/internal/assistant/lifecyclepayload/lifecyclepayload_test.go +++ b/internal/assistant/lifecyclepayload/lifecyclepayload_test.go @@ -34,6 +34,7 @@ func TestPromptAndTurnPayloads(t *testing.T) { parentID := lifecycleTestParentID prompt := lifecyclepayload.Prompt(&lifecyclepayload.PromptRequest{ ParentEntryID: &parentID, + Attachments: nil, CWD: "/work", Name: "agent", SessionID: lifecycleTestSessionID, @@ -109,7 +110,7 @@ func TestSessionEntryAndContextPayloads(t *testing.T) { Role: database.RoleAssistant, Content: "answer", Provider: "provider-1", - Model: lifecycleTestModel, + Model: lifecycleTestModel, Parts: nil, }, Summary: "summary", ToolStatus: "", diff --git a/internal/assistant/llm_conversion.go b/internal/assistant/llm_conversion.go index acbc2970..0ca7d1cb 100644 --- a/internal/assistant/llm_conversion.go +++ b/internal/assistant/llm_conversion.go @@ -1,6 +1,7 @@ package assistant import ( + "encoding/base64" "strings" "github.com/samber/lo" @@ -47,7 +48,7 @@ func llmMessagesFromDatabase(messages []database.MessageEntity) []llm.Message { } func llmMessageFromDatabase(message *database.MessageEntity) (llm.Message, bool) { - if message == nil || strings.TrimSpace(message.Content) == "" { + if message == nil { return emptyLLMMessage(), false } @@ -56,7 +57,37 @@ func llmMessageFromDatabase(message *database.MessageEntity) (llm.Message, bool) return emptyLLMMessage(), false } - return llm.TextMessage(role, message.Content), true + parts := llmPartsFromDatabase(message.Parts) + + if len(parts) == 0 && strings.TrimSpace(message.Content) != "" { + parts = append(parts, llm.TextPart(message.Content)) + } + + if len(parts) == 0 { + return emptyLLMMessage(), false + } + + return llm.Message{Metadata: nil, Role: role, Content: parts}, true +} + +func llmPartsFromDatabase(databaseParts []database.MessagePartEntity) []llm.Part { + parts := make([]llm.Part, 0, len(databaseParts)) + for index := range databaseParts { + part := &databaseParts[index] + if part.Type == database.MessagePartText && strings.TrimSpace(part.Text) != "" { + parts = append(parts, llm.TextPart(part.Text)) + } + + if part.Type == database.MessagePartImage && len(part.Data) != 0 { + parts = append(parts, llm.Part{ + Metadata: imagePartMetadata(part.Name, part.Width, part.Height), + ToolCall: nil, ToolResult: nil, Type: llm.PartImage, Text: "", + Data: base64.StdEncoding.EncodeToString(part.Data), MIMEType: part.MIMEType, + }) + } + } + + return parts } func emptyLLMMessage() llm.Message { diff --git a/internal/assistant/llm_conversion_internal_test.go b/internal/assistant/llm_conversion_internal_test.go index 726b9c15..58d22139 100644 --- a/internal/assistant/llm_conversion_internal_test.go +++ b/internal/assistant/llm_conversion_internal_test.go @@ -1,7 +1,9 @@ package assistant import ( + "encoding/base64" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -85,6 +87,38 @@ func TestLLMRequestFromCompletionRequestConvertsAssistantState(t *testing.T) { assert.Equal(t, "yes", request.Model.Compat["compat"]) } +func TestLLMMessageFromDatabasePreservesOrderedMultipartImageOnly(t *testing.T) { + t.Parallel() + + data := []byte{0, 1, 2, 3} + message, converted := llmMessageFromDatabase(&database.MessageEntity{ + Timestamp: time.Time{}, Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{ + {Text: "inspect image", MIMEType: "", Name: "", Type: database.MessagePartText, + Data: nil, Width: 0, Height: 0}, + {Text: "", MIMEType: imageMIMEPNG, Name: "screen.png", Type: database.MessagePartImage, + Data: data, Width: 10, Height: 20}, + }, + }) + + require.True(t, converted) + require.Len(t, message.Content, 2) + assert.Equal(t, llm.PartText, message.Content[0].Type) + assert.Equal(t, llm.PartImage, message.Content[1].Type) + assert.Equal(t, base64.StdEncoding.EncodeToString(data), message.Content[1].Data) + assert.Equal(t, "screen.png", message.Content[1].Metadata["name"]) + + imageOnly, imageOnlyConverted := llmMessageFromDatabase(&database.MessageEntity{ + Timestamp: time.Time{}, Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{{ + Text: "", MIMEType: imageMIMEPNG, Name: "", Type: database.MessagePartImage, + Data: data, Width: 1, Height: 1, + }}, + }) + require.True(t, imageOnlyConverted) + assert.Len(t, imageOnly.Content, 1) +} + func TestLLMRequestFromCompletionRequestNilAndDisabledTools(t *testing.T) { t.Parallel() diff --git a/internal/assistant/prompt_images.go b/internal/assistant/prompt_images.go new file mode 100644 index 00000000..3f379391 --- /dev/null +++ b/internal/assistant/prompt_images.go @@ -0,0 +1,231 @@ +package assistant + +import ( + "bytes" + "image" + _ "image/gif" // Register supported image decoders for DecodeConfig. + _ "image/jpeg" // Register supported image decoders for DecodeConfig. + _ "image/png" // Register supported image decoders for DecodeConfig. + "slices" + "strings" + + "github.com/samber/oops" + + "github.com/omarluq/librecode/internal/database" + "github.com/omarluq/librecode/internal/model" +) + +const ( + maxPromptImages = 4 + maxPromptImageName = 255 + maxPromptImageBytes = 5 << 20 + maxPromptImageTotal = 20 << 20 + maxPromptImagePixels = 40_000_000 + imageMIMEPNG = "image/png" +) + +func (runtime *Runtime) preparePromptRequest(request *PromptRequest) (*PromptRequest, error) { + if err := validatePromptImageAllocationBounds(request.Images); err != nil { + return nil, err + } + + cloned := clonePromptRequest(request) + if err := runtime.validatePromptRequest(cloned); err != nil { + return nil, err + } + + return cloned, nil +} + +func clonePromptRequest(request *PromptRequest) *PromptRequest { + cloned := *request + + cloned.Images = make([]ImageAttachment, len(request.Images)) + for index := range request.Images { + cloned.Images[index] = request.Images[index] + cloned.Images[index].Data = bytes.Clone(request.Images[index].Data) + } + + return &cloned +} + +func (runtime *Runtime) validatePromptRequest(request *PromptRequest) error { + if strings.TrimSpace(request.Text) == "" && len(request.Images) == 0 { + return oops.In("assistant").Code("empty_prompt").Errorf("prompt text and images are empty") + } + + if len(request.Images) == 0 { + return nil + } + + if strings.HasPrefix(strings.TrimSpace(request.Text), slashPrefix) { + return oops.In("assistant").Code("slash_command_images"). + Errorf("slash commands do not accept image attachments") + } + + if err := validatePromptImages(request.Images); err != nil { + return err + } + + if runtime.models == nil { + return oops.In("assistant").Code("models_unavailable").Errorf("model registry is not configured") + } + + selected, err := runtime.selectedModel() + if err != nil { + return err + } + + return validateSelectedModelImageInput(&selected, request.Images) +} + +func validatePromptImageAllocationBounds(images []ImageAttachment) error { + if len(images) > maxPromptImages { + return promptImageError( + "image_count_limit", "prompt has %d images; maximum is %d", len(images), maxPromptImages, + ) + } + + total := 0 + + for index := range images { + if len(images[index].Name) > maxPromptImageName { + return promptImageError( + "image_name_limit", "image %d name exceeds %d bytes", index+1, maxPromptImageName, + ) + } + + if len(images[index].Data) > maxPromptImageBytes { + return promptImageError("image_size_limit", "image %d exceeds the 5 MiB limit", index+1) + } + + total += len(images[index].Data) + if total > maxPromptImageTotal { + return promptImageError("image_total_size_limit", "image attachments exceed the 20 MiB limit") + } + } + + return nil +} + +func validatePromptImages(images []ImageAttachment) error { + if err := validatePromptImageAllocationBounds(images); err != nil { + return err + } + + for index := range images { + if err := validatePromptImage(&images[index], index); err != nil { + return err + } + } + + return nil +} + +func validatePromptImage(attachment *ImageAttachment, index int) error { + if len(attachment.Data) == 0 { + return promptImageError("invalid_image", "image %d is empty", index+1) + } + + if len(attachment.Data) > maxPromptImageBytes { + return promptImageError("image_size_limit", "image %d exceeds the 5 MiB limit", index+1) + } + + config, format, err := image.DecodeConfig(bytes.NewReader(attachment.Data)) + if err != nil { + return oops.In("assistant").Code("invalid_image"). + With("image_index", index).Wrapf(err, "decode image %d", index+1) + } + + return validatePromptImageMetadata(attachment, config, format, index) +} + +func validatePromptImageMetadata(attachment *ImageAttachment, config image.Config, format string, index int) error { + mimeType := imageMIMEType(format) + if mimeType == "" || attachment.MIMEType != mimeType { + return promptImageError( + "invalid_image_mime", "image %d MIME type %q does not match format %q", + index+1, attachment.MIMEType, format, + ) + } + + if config.Width <= 0 || config.Height <= 0 || config.Width > maxPromptImagePixels/config.Height { + return promptImageError( + "image_dimensions_limit", "image %d dimensions must be positive and at most 40 megapixels", index+1, + ) + } + + if attachment.Width != config.Width || attachment.Height != config.Height { + return promptImageError("invalid_image_dimensions", "image %d dimensions do not match its data", index+1) + } + + return nil +} + +func imageMIMEType(format string) string { + switch format { + case "gif": + return "image/gif" + case "jpeg": + return "image/jpeg" + case "png": + return imageMIMEPNG + default: + return "" + } +} + +func promptImageError(code, format string, args ...any) error { + return oops.In("assistant").Code(code).Errorf(format, args...) +} + +func imageMetadata(attachment ImageAttachment) map[string]any { + return map[string]any{ + executeNameKey: attachment.Name, "mime_type": attachment.MIMEType, + "width": attachment.Width, "height": attachment.Height, "size": len(attachment.Data), + } +} + +func validateSelectedModelImageInput(selected *model.Model, images []ImageAttachment) error { + if len(images) == 0 { + return nil + } + + return validateSelectedModelHasImageInput(selected, "prompt") +} + +func validateSelectedModelHasImageInput(selected *model.Model, source string) error { + if slices.Contains(selected.Input, model.InputImage) { + return nil + } + + return oops.In("assistant").Code("image_input_unsupported"). + With("image_source", source). + With("provider", selected.Provider). + With("model", selected.ID). + Errorf("selected model %s/%s does not support image input", selected.Provider, selected.ID) +} + +func messagesContainImages(messages []database.MessageEntity) bool { + for messageIndex := range messages { + for partIndex := range messages[messageIndex].Parts { + if messages[messageIndex].Parts[partIndex].Type == database.MessagePartImage { + return true + } + } + } + + return false +} + +func validateModelContextImageInput(selected *model.Model, messages []database.MessageEntity) error { + if !messagesContainImages(messages) || slices.Contains(selected.Input, model.InputImage) { + return nil + } + + return validateSelectedModelHasImageInput(selected, "conversation_history") +} + +func imagePartMetadata(name string, width, height int) map[string]any { + return map[string]any{executeNameKey: name, "width": width, "height": height} +} diff --git a/internal/assistant/prompt_images_internal_test.go b/internal/assistant/prompt_images_internal_test.go new file mode 100644 index 00000000..f8ec35ee --- /dev/null +++ b/internal/assistant/prompt_images_internal_test.go @@ -0,0 +1,112 @@ +package assistant + +import ( + "bytes" + "image" + "image/png" + "testing" + "time" + + "github.com/samber/oops" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/omarluq/librecode/internal/database" + "github.com/omarluq/librecode/internal/model" +) + +func testPNG(t *testing.T, width, height int) []byte { + t.Helper() + + var output bytes.Buffer + require.NoError(t, png.Encode(&output, image.NewRGBA(image.Rect(0, 0, width, height)))) + + return output.Bytes() +} + +func TestValidatePromptImagesAndCloneBoundary(t *testing.T) { + t.Parallel() + + data := testPNG(t, 2, 3) + images := []ImageAttachment{{Name: "screen.png", MIMEType: imageMIMEPNG, Data: data, Width: 2, Height: 3}} + require.NoError(t, validatePromptImages(images)) + + request := &PromptRequest{ + OnEvent: nil, OnRetry: nil, OnUserEntry: nil, ParentEntryID: nil, + SessionID: "", CWD: "", Text: "", Images: images, Name: "", ResumeLatest: false, HideUserPrompt: false, + } + cloned := clonePromptRequest(request) + cloned.Images[0].Data[0]++ + assert.NotEqual(t, cloned.Images[0].Data[0], request.Images[0].Data[0]) +} + +func TestValidateSelectedModelImageInput(t *testing.T) { + t.Parallel() + + images := []ImageAttachment{{ + Name: "", MIMEType: imageMIMEPNG, Data: []byte{1}, Width: 1, Height: 1, + }} + vision := promptImageTestModel("vision", []model.InputMode{model.InputText, model.InputImage}) + require.NoError(t, validateSelectedModelImageInput(&vision, images)) + + textOnly := promptImageTestModel("text-only", []model.InputMode{model.InputText}) + err := validateSelectedModelImageInput(&textOnly, images) + require.Error(t, err) + coded, ok := oops.AsOops(err) + require.True(t, ok) + assert.Equal(t, "image_input_unsupported", coded.Code()) +} + +func promptImageTestModel(id string, input []model.InputMode) model.Model { + return model.Model{ + ThinkingLevelMap: nil, Headers: nil, Compat: nil, Provider: "image-test-provider", ID: id, + Name: id, API: "", BaseURL: "", Input: input, + Cost: model.Cost{Input: 0, Output: 0, CacheRead: 0, CacheWrite: 0}, + ContextWindow: 0, MaxTokens: 0, Reasoning: false, + } +} + +func TestValidatePromptImagesRejectsUnsafeInput(t *testing.T) { + t.Parallel() + + data := testPNG(t, 2, 3) + + tests := []ImageAttachment{ + {Name: "", MIMEType: "image/jpeg", Data: data, Width: 2, Height: 3}, + {Name: "", MIMEType: imageMIMEPNG, Data: data, Width: 3, Height: 2}, + {Name: "", MIMEType: imageMIMEPNG, Data: []byte("not an image"), Width: 1, Height: 1}, + } + for _, attachment := range tests { + require.Error(t, validatePromptImages([]ImageAttachment{attachment})) + } + + many := make([]ImageAttachment, maxPromptImages+1) + require.Error(t, validatePromptImages(many)) + + longName := ImageAttachment{ + Name: string(bytes.Repeat([]byte{'x'}, maxPromptImageName+1)), + MIMEType: imageMIMEPNG, Data: data, Width: 2, Height: 3, + } + require.Error(t, validatePromptImages([]ImageAttachment{longName})) +} + +func TestValidateModelContextImageInput(t *testing.T) { + t.Parallel() + + messages := []database.MessageEntity{{ + Timestamp: time.Time{}, Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{{ + Text: "", MIMEType: imageMIMEPNG, Name: "", Type: database.MessagePartImage, + Data: []byte{1}, Width: 1, Height: 1, + }}, + }} + vision := promptImageTestModel("vision", []model.InputMode{model.InputText, model.InputImage}) + require.NoError(t, validateModelContextImageInput(&vision, messages)) + + textOnly := promptImageTestModel("text-only", []model.InputMode{model.InputText}) + err := validateModelContextImageInput(&textOnly, messages) + require.Error(t, err) + coded, ok := oops.AsOops(err) + require.True(t, ok) + assert.Equal(t, "image_input_unsupported", coded.Code()) +} diff --git a/internal/assistant/runtime.go b/internal/assistant/runtime.go index 837e1164..46685581 100644 --- a/internal/assistant/runtime.go +++ b/internal/assistant/runtime.go @@ -45,6 +45,15 @@ type Runtime struct { profile ExecutionProfile } +// ImageAttachment is one provider-neutral image supplied with a prompt. +type ImageAttachment struct { + Name string `json:"name,omitempty"` + MIMEType string `json:"mime_type"` + Data []byte `json:"-"` + Width int `json:"width"` + Height int `json:"height"` +} + // PromptRequest contains one user prompt invocation. type PromptRequest struct { OnEvent func(StreamEvent) `json:"-"` @@ -55,6 +64,7 @@ type PromptRequest struct { CWD string `json:"cwd"` Text string `json:"text"` Name string `json:"name"` + Images []ImageAttachment `json:"images,omitempty"` ResumeLatest bool `json:"resume_latest,omitempty"` HideUserPrompt bool `json:"-"` } @@ -197,13 +207,14 @@ func (runtime *Runtime) Prompt(ctx context.Context, request *PromptRequest) (res return nil, oops.In("assistant").Code("nil_prompt_request").Errorf("prompt request is nil") } - promptPayload := lifecyclePromptRequest(request) - runtime.dispatchObservationalLifecycle(ctx, extension.LifecycleInput, lifecyclepayload.Prompt(promptPayload)) - runtime.dispatchObservationalLifecycle( - ctx, - extension.LifecyclePromptPrepare, - lifecyclepayload.Prompt(promptPayload), - ) + originalRequest := request + + request, err = runtime.preparePromptRequest(request) + if err != nil { + return nil, err + } + + runtime.dispatchPromptInputLifecycle(ctx, request) persistCtx, persistCancel := promptPersistenceContext(ctx) defer persistCancel() @@ -213,6 +224,8 @@ func (runtime *Runtime) Prompt(ctx context.Context, request *PromptRequest) (res return nil, err } + originalRequest.SessionID = activeSession.ID + releaseOperation, err := runtime.acquirePromptOperation(ctx, activeSession.ID) if err != nil { return nil, err @@ -270,6 +283,12 @@ func (runtime *Runtime) Prompt(ctx context.Context, request *PromptRequest) (res }, nil } +func (runtime *Runtime) dispatchPromptInputLifecycle(ctx context.Context, request *PromptRequest) { + payload := lifecyclepayload.Prompt(lifecyclePromptRequest(request)) + runtime.dispatchObservationalLifecycle(ctx, extension.LifecycleInput, payload) + runtime.dispatchObservationalLifecycle(ctx, extension.LifecyclePromptPrepare, payload) +} + func (runtime *Runtime) persistAssistantBundle( ctx context.Context, sessionID string, @@ -304,6 +323,7 @@ func (runtime *Runtime) appendPromptUserEntry( activeSession.ID, parentID, request.Text, + request.Images, !request.HideUserPrompt, ) if err != nil { diff --git a/internal/assistant/runtime_context_internal_test.go b/internal/assistant/runtime_context_internal_test.go index e2791869..b1d1645a 100644 --- a/internal/assistant/runtime_context_internal_test.go +++ b/internal/assistant/runtime_context_internal_test.go @@ -134,7 +134,7 @@ func newRuntimeContextTestMessage(role database.Role, content string) database.M Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, } } diff --git a/internal/assistant/runtime_entries.go b/internal/assistant/runtime_entries.go index 34b08095..a0767c7a 100644 --- a/internal/assistant/runtime_entries.go +++ b/internal/assistant/runtime_entries.go @@ -3,6 +3,7 @@ package assistant import ( "context" + "strings" "time" "github.com/omarluq/librecode/internal/contextwindow" @@ -14,16 +15,34 @@ func (runtime *Runtime) appendUserPromptEntry( sessionID string, parentID *string, prompt string, + images []ImageAttachment, display bool, ) (*database.EntryEntity, error) { + parts := make([]database.MessagePartEntity, 0, len(images)+1) + + content := prompt + if strings.TrimSpace(prompt) != "" { + parts = append(parts, database.MessagePartEntity{ + Text: prompt, MIMEType: "", Name: "", Type: database.MessagePartText, + Data: nil, Width: 0, Height: 0, + }) + } else { + content = "" + } + + for index := range images { + image := &images[index] + parts = append(parts, database.MessagePartEntity{ + Text: "", MIMEType: image.MIMEType, Name: image.Name, Type: database.MessagePartImage, + Data: append([]byte(nil), image.Data...), Width: image.Width, Height: image.Height, + }) + } + message := database.MessageEntity{ - Timestamp: time.Now().UTC(), - Role: database.RoleUser, - Content: prompt, - Provider: "", - Model: "", + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: content, + Provider: "", Model: "", Parts: parts, } - modelFacing := promptModelFacing(prompt) + modelFacing := len(images) > 0 || promptModelFacing(prompt) entry, err := runtime.sessions.AppendMessageWithDisplay( ctx, @@ -48,7 +67,7 @@ func (runtime *Runtime) appendAssistantResponseEntry( Role: database.RoleAssistant, Content: bundle.Text, Provider: runtime.cfg.Assistant.Provider, - Model: runtime.cfg.Assistant.Model, + Model: runtime.cfg.Assistant.Model, Parts: nil, } entry, err := runtime.sessions.AppendMessageWithMetadata( diff --git a/internal/assistant/runtime_model.go b/internal/assistant/runtime_model.go index 83b24d84..19c4eeaf 100644 --- a/internal/assistant/runtime_model.go +++ b/internal/assistant/runtime_model.go @@ -37,6 +37,7 @@ func (runtime *Runtime) respond( lineage *promptLineage, cwd string, prompt string, + hasImages bool, onEvent func(StreamEvent), onRetry RetryEventHandler, ) ( @@ -44,7 +45,7 @@ func (runtime *Runtime) respond( cached bool, err error, ) { - if strings.HasPrefix(prompt, slashPrefix) { + if strings.HasPrefix(strings.TrimSpace(prompt), slashPrefix) { slashResponse, slashToolEvents, slashErr := runtime.respondToSlashCommand(ctx, cwd, prompt, onEvent) return &responseBundle{ @@ -58,6 +59,21 @@ func (runtime *Runtime) respond( cacheKey := runtime.cacheKey(sessionID, prompt) + contextHasImages, contextErr := runtime.promptContextContainsImages(ctx, sessionID, lineage) + if contextErr != nil { + return nil, false, contextErr + } + + if hasImages || contextHasImages { + // Image bytes deliberately stay out of cache keys; prompts with image-bearing + // context always execute against their durable multipart history. + imageBundle, modelErr := runtime.modelResponse( + ctx, sessionID, lineage, cwd, prompt, contextHasImages, onEvent, onRetry, + ) + + return imageBundle, false, modelErr + } + cachedResponse, found, err := runtime.cache.Get(cacheKey) if err != nil { return nil, false, oops.In("assistant").Code("cache_get").Wrapf(err, "read response cache") @@ -73,7 +89,7 @@ func (runtime *Runtime) respond( }, true, nil } - bundle, err = runtime.modelResponse(ctx, sessionID, lineage, cwd, prompt, onEvent, onRetry) + bundle, err = runtime.modelResponse(ctx, sessionID, lineage, cwd, prompt, false, onEvent, onRetry) if err != nil { return nil, false, err } @@ -89,6 +105,7 @@ func (runtime *Runtime) modelResponse( lineage *promptLineage, cwd string, prompt string, + contextHasImages bool, onEvent func(StreamEvent), onRetry RetryEventHandler, ) (*responseBundle, error) { @@ -101,6 +118,13 @@ func (runtime *Runtime) modelResponse( return nil, err } + if contextHasImages { + imageErr := validateSelectedModelHasImageInput(&selectedModel, "conversation_history") + if imageErr != nil { + return nil, imageErr + } + } + auth := runtime.models.RequestAuthContext(ctx, selectedModel.Provider) if !auth.OK { return nil, oops.In("assistant"). @@ -124,6 +148,13 @@ func (runtime *Runtime) modelResponse( return nil, err } + if contextImageErr := validateModelContextImageInput( + &selectedModel, + build.Request.Messages, + ); contextImageErr != nil { + return nil, contextImageErr + } + build, compactionEntry, result, err := runtime.completeWithProviderOverflowRecovery( ctx, &providerOverflowRecoveryInput{ @@ -389,6 +420,23 @@ func (runtime *Runtime) thinkingLevel() string { return runtime.cfg.Assistant.ThinkingLevel } +func (runtime *Runtime) promptContextContainsImages( + ctx context.Context, + sessionID string, + lineage *promptLineage, +) (bool, error) { + if lineage == nil || strings.TrimSpace(lineage.activeParentEntryID) == "" { + return false, nil + } + + contextEntity, err := runtime.modelContextEntityFrom(ctx, sessionID, lineage.activeParentEntryID) + if err != nil { + return false, err + } + + return messagesContainImages(contextEntity.Messages), nil +} + func (runtime *Runtime) cacheKey(sessionID, prompt string) string { selected, err := runtime.selectedModel() if err != nil { diff --git a/internal/assistant/runtime_persist.go b/internal/assistant/runtime_persist.go index 066c199d..c45e2dc4 100644 --- a/internal/assistant/runtime_persist.go +++ b/internal/assistant/runtime_persist.go @@ -64,7 +64,7 @@ func (runtime *Runtime) appendAssistantSideEffects( Role: database.RoleThinking, Content: trimmed, Provider: runtime.cfg.Assistant.Provider, - Model: runtime.cfg.Assistant.Model, + Model: runtime.cfg.Assistant.Model, Parts: nil, } entry, err := runtime.sessions.AppendMessage(ctx, sessionID, parentID, &message) @@ -83,7 +83,7 @@ func (runtime *Runtime) appendAssistantSideEffects( Role: database.RoleToolResult, Content: formatToolEvent(event), Provider: runtime.cfg.Assistant.Provider, - Model: runtime.cfg.Assistant.Model, + Model: runtime.cfg.Assistant.Model, Parts: nil, } entry, err := runtime.sessions.AppendMessage(ctx, sessionID, parentID, &message) @@ -112,6 +112,7 @@ func (runtime *Runtime) respondWithPartialProgress( lineage, request.CWD, request.Text, + len(request.Images) > 0, progress.handle, progress.retryHandler(request.OnRetry), ) @@ -384,7 +385,7 @@ func (runtime *Runtime) appendPartialPromptMessages( Role: partial.Role, Content: partial.Content, Provider: runtime.cfg.Assistant.Provider, - Model: runtime.cfg.Assistant.Model, + Model: runtime.cfg.Assistant.Model, Parts: nil, } entry, err := runtime.sessions.AppendMessage(ctx, sessionID, parentID, &message) diff --git a/internal/assistant/runtime_test.go b/internal/assistant/runtime_test.go index 88986170..03ec94bf 100644 --- a/internal/assistant/runtime_test.go +++ b/internal/assistant/runtime_test.go @@ -1,9 +1,12 @@ package assistant_test import ( + "bytes" "context" "database/sql" "errors" + "image" + "image/png" "io" "log/slog" "os" @@ -56,6 +59,7 @@ func newRuntimePromptRequest(cwd, text, name string) *assistant.PromptRequest { ParentEntryID: nil, SessionID: "", CWD: cwd, + Images: nil, Text: text, Name: name, ResumeLatest: false, @@ -126,6 +130,75 @@ func TestRuntime_PromptPersistsConversation(t *testing.T) { assert.Equal(t, database.RoleAssistant, entries[1].Message.Role) } +func TestRuntime_PromptImagesRoundTripAndDisableResponseCache(t *testing.T) { + t.Parallel() + + client := &countingCompleter{request: nil, attempts: 0} + runtime, repository := newTestRuntimeWithClient(t, client) + imageData := runtimeTestPNG(t, 2, 3) + request := newRuntimePromptRequest(testRuntimeCWD, " ", "images") + request.Images = []assistant.ImageAttachment{{ + Name: "screen.png", MIMEType: "image/png", Data: imageData, Width: 2, Height: 3, + }} + + response, err := runtime.Prompt(t.Context(), request) + require.NoError(t, err) + assert.Equal(t, 1, client.attempts) + + messages := requireRuntimeMessages(t, repository, response.SessionID, 2) + require.Len(t, messages[0].Parts, 1) + assert.Equal(t, database.MessagePartImage, messages[0].Parts[0].Type) + assert.Equal(t, imageData, messages[0].Parts[0].Data) + require.NotNil(t, client.request) + require.Len(t, client.request.Messages[0].Parts, 1) + assert.Equal(t, database.MessagePartImage, client.request.Messages[0].Parts[0].Type) + + for range 2 { + followUp := newRuntimePromptRequest(testRuntimeCWD, "same follow-up", "") + followUp.SessionID = response.SessionID + _, err = runtime.Prompt(t.Context(), followUp) + require.NoError(t, err) + } + + assert.Equal(t, 3, client.attempts, "image-bearing history must bypass the text response cache") +} + +func TestRuntime_TextOnlyModelRejectsImageHistoryBeforeProviderCall(t *testing.T) { + t.Parallel() + + imageClient := &countingCompleter{request: nil, attempts: 0} + imageRuntime, repository := newTestRuntimeWithClient(t, imageClient) + request := newRuntimePromptRequest(testRuntimeCWD, "look", "images") + request.Images = []assistant.ImageAttachment{{ + Name: "screen.png", MIMEType: "image/png", Data: runtimeTestPNG(t, 1, 1), Width: 1, Height: 1, + }} + response, err := imageRuntime.Prompt(t.Context(), request) + require.NoError(t, err) + + textClient := &countingCompleter{request: nil, attempts: 0} + textRuntime, _ := newTestRuntimeWithRepositoryClientAndInput( + t, repository, textClient, []model.InputMode{model.InputText}, + ) + followUp := newRuntimePromptRequest(testRuntimeCWD, "continue", "") + followUp.SessionID = response.SessionID + _, err = textRuntime.Prompt(t.Context(), followUp) + require.Error(t, err) + assert.Equal(t, 0, textClient.attempts) + + coded, ok := oops.AsOops(err) + require.True(t, ok) + assert.Equal(t, "image_input_unsupported", coded.Code()) +} + +func runtimeTestPNG(t *testing.T, width, height int) []byte { + t.Helper() + + var output bytes.Buffer + require.NoError(t, png.Encode(&output, image.NewRGBA(image.Rect(0, 0, width, height)))) + + return output.Bytes() +} + func TestRuntime_HiddenPromptDoesNotAppearInTranscript(t *testing.T) { t.Parallel() @@ -558,7 +631,7 @@ func TestRuntime_PromptEstimatesContextFromModelFacingBranch(t *testing.T) { Role: database.RoleUser, Content: "hello", Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) _, err = repository.AppendMessage(ctx, session.ID, &userEntry.ID, &database.MessageEntity{ @@ -566,7 +639,7 @@ func TestRuntime_PromptEstimatesContextFromModelFacingBranch(t *testing.T) { Role: database.RoleToolResult, Content: strings.Repeat("tool output ", 10_000), Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) @@ -608,7 +681,7 @@ func TestRuntime_PromptIncludesCompactionSummaryContext(t *testing.T) { Role: database.RoleUser, Content: "old user prompt", Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) compactionEntry, err := repository.AppendCompaction(ctx, &database.AppendCompactionInput{ @@ -724,7 +797,34 @@ func newTestRuntimeWithRepositoryAndClient( ) (*assistant.Runtime, *database.SessionRepository) { t.Helper() - runtime, repository, _ := newTestRuntimeWithRepositoryClientAndManager(t, repository, client) + return newTestRuntimeWithRepositoryClientAndInput( + t, repository, client, []model.InputMode{model.InputText, model.InputImage}, + ) +} + +func newTestRuntimeWithRepositoryClientAndInput( + t *testing.T, + repository *database.SessionRepository, + client assistant.Completer, + input []model.InputMode, +) (*assistant.Runtime, *database.SessionRepository) { + t.Helper() + + manager := extension.NewManager(slog.New(slog.NewTextHandler(io.Discard, nil))) + t.Cleanup(manager.Shutdown) + + cache := assistant.NewResponseCache(true, 32, time.Minute) + t.Cleanup(cache.Shutdown) + + runtime := assistant.NewRuntimeForTest(func(opts *assistant.RuntimeTestOptions) { + opts.Config = testConfig() + opts.Sessions = repository + opts.Extensions = manager + opts.Cache = cache + opts.Models = testRegistryWithInput(t, input) + opts.Client = client + opts.Logger = slog.New(slog.NewTextHandler(io.Discard, nil)) + }) return runtime, repository } @@ -773,6 +873,11 @@ type capturingCompleter struct { request *assistant.CompletionRequest } +type countingCompleter struct { + request *assistant.CompletionRequest + attempts int +} + type retryCompleter struct { err error response string @@ -816,6 +921,16 @@ func (client *capturingCompleter) Complete( return testCompleter{}.Complete(ctx, request) } +func (client *countingCompleter) Complete( + ctx context.Context, + request *assistant.CompletionRequest, +) (*assistant.CompletionResult, error) { + client.request = request + client.attempts++ + + return testCompleter{}.Complete(ctx, request) +} + func (client *retryCompleter) Complete( _ context.Context, request *assistant.CompletionRequest, @@ -1098,6 +1213,12 @@ func (testCompleter) Complete( func testRegistry(t *testing.T) *model.Registry { t.Helper() + return testRegistryWithInput(t, []model.InputMode{model.InputText, model.InputImage}) +} + +func testRegistryWithInput(t *testing.T, input []model.InputMode) *model.Registry { + t.Helper() + storage := testutil.NewAuthStorage(t, map[string]auth.Credential{ testRuntimeProvider: testProviderCredential(), }) @@ -1116,7 +1237,7 @@ func testRegistry(t *testing.T) *model.Registry { Name: testRuntimeModel, API: "openai-completions", BaseURL: "https://example.invalid/v1", - Input: []model.InputMode{model.InputText}, + Input: input, Cost: model.Cost{Input: 0, Output: 0, CacheRead: 0, CacheWrite: 0}, ContextWindow: 100_000, MaxTokens: 0, diff --git a/internal/assistant/test_message_helpers_internal_test.go b/internal/assistant/test_message_helpers_internal_test.go index 526fcc0f..766663f8 100644 --- a/internal/assistant/test_message_helpers_internal_test.go +++ b/internal/assistant/test_message_helpers_internal_test.go @@ -12,6 +12,6 @@ func testMessageEntity(role database.Role, content string) database.MessageEntit Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, } } diff --git a/internal/compaction/plan.go b/internal/compaction/plan.go index ca9f8318..68953eae 100644 --- a/internal/compaction/plan.go +++ b/internal/compaction/plan.go @@ -9,6 +9,7 @@ import ( "github.com/samber/oops" + "github.com/omarluq/librecode/internal/contextwindow" "github.com/omarluq/librecode/internal/database" "github.com/omarluq/librecode/internal/model" ) @@ -377,9 +378,10 @@ func formatSplitTurnSummary(messages []database.MessageEntity) string { for index := range messages { message := messages[index] - content := strings.TrimSpace(message.Content) - if content == "" { + + imageDescription := messageImageDescription(&message) + if content == "" && imageDescription == "" { continue } @@ -387,6 +389,12 @@ func formatSplitTurnSummary(messages []database.MessageEntity) string { builder.WriteString(string(message.Role)) builder.WriteString("]\n") builder.WriteString(content) + + if content != "" && imageDescription != "" { + builder.WriteString("\n") + } + + builder.WriteString(imageDescription) builder.WriteString("\n") } @@ -484,7 +492,7 @@ func emptyMessage() database.MessageEntity { Role: "", Content: "", Provider: "", - Model: "", + Model: "", Parts: nil, } } @@ -528,13 +536,45 @@ func candidateMessage(entry *database.EntryEntity) database.MessageEntity { func messageTokens(messages []database.MessageEntity, countTokens TokenCounter) int { tokens := 0 + for index := range messages { - tokens += countTokens(messages[index].Content) + message := &messages[index] + tokens += countTokens(message.Content) + + imageParts := make([]database.MessagePartEntity, 0, len(message.Parts)) + for partIndex := range message.Parts { + if message.Parts[partIndex].Type == database.MessagePartImage { + imageParts = append(imageParts, message.Parts[partIndex]) + } + } + + if len(imageParts) > 0 { + imageOnly := database.MessageEntity{ + Timestamp: time.Time{}, Role: "", Content: "", Provider: "", Model: "", Parts: imageParts, + } + tokens += contextwindow.EstimateMessageTokens([]database.MessageEntity{imageOnly}) + } } return tokens } +func messageImageDescription(message *database.MessageEntity) string { + images := 0 + + for index := range message.Parts { + if message.Parts[index].Type == database.MessagePartImage { + images++ + } + } + + if images == 0 { + return "" + } + + return fmt.Sprintf("[attached images: %d]", images) +} + func countRunesAsTokens(text string) int { return utf8.RuneCountInString(strings.TrimSpace(text)) } diff --git a/internal/compaction/plan_internal_test.go b/internal/compaction/plan_internal_test.go index e698c428..05d2ad53 100644 --- a/internal/compaction/plan_internal_test.go +++ b/internal/compaction/plan_internal_test.go @@ -11,6 +11,54 @@ import ( "github.com/omarluq/librecode/internal/database" ) +func TestImageMessagesContributeToPlanningAndSplitSummary(t *testing.T) { + t.Parallel() + + imagePart := database.MessagePartEntity{ + Text: "", MIMEType: "image/png", Name: "", Type: database.MessagePartImage, + Data: []byte{1}, Width: 512, Height: 512, + } + + tests := []struct { + name string + content string + parts []database.MessagePartEntity + }{ + {name: "image only", content: "", parts: []database.MessagePartEntity{imagePart}}, + { + name: "text and image", content: "describe this", + parts: []database.MessagePartEntity{ + { + Text: "describe this", MIMEType: "", Name: "", Type: database.MessagePartText, + Data: nil, Width: 0, Height: 0, + }, + imagePart, + }, + }, + } + + for _, testCase := range tests { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + message := database.MessageEntity{ + Timestamp: time.Time{}, Role: database.RoleUser, Content: testCase.content, + Provider: "", Model: "", Parts: testCase.parts, + } + entry := database.EntryEntity{ + CreatedAt: time.Time{}, ParentID: nil, ToolStatus: "", SessionID: "", ToolArgsJSON: "", + CustomType: "", DataJSON: "", ID: "", Summary: "", ToolName: "", BranchFromEntryID: "", + CompactionFirstKeptEntryID: "", Message: message, Type: database.EntryTypeMessage, + CompactionTokensBefore: 0, TokenEstimate: 0, Display: false, ModelFacing: true, + } + + wantTokens := estimateTokens(testCase.content) + 256 + assert.Equal(t, wantTokens, BranchTokens([]database.EntryEntity{entry}, estimateTokens)) + assert.Contains(t, formatSplitTurnSummary([]database.MessageEntity{message}), "[attached images: 1]") + }) + } +} + type planCompactionCase struct { assertFn func(t *testing.T, plan *Plan) name string @@ -282,7 +330,7 @@ func testEntry( Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, }, Summary: "", ToolStatus: "", diff --git a/internal/contextwindow/tokens.go b/internal/contextwindow/tokens.go index ffff4cca..b43c6c9e 100644 --- a/internal/contextwindow/tokens.go +++ b/internal/contextwindow/tokens.go @@ -7,7 +7,13 @@ import ( "github.com/omarluq/librecode/internal/database" ) -const charsPerEstimatedToken = 4 +const ( + charsPerEstimatedToken = 4 + imageTilePixels = 512 + imageTileTokens = 200 + minimumImageTokens = 256 + maximumImageTokens = 16_000 +) // EstimateTokens returns a rough cross-provider estimate used until provider usage arrives. func EstimateTokens(text string) int { @@ -28,7 +34,7 @@ func EstimateTokens(text string) int { func EstimateInputTokens(systemPrompt string, messages []database.MessageEntity) int { count := EstimateTokens(systemPrompt) for index := range messages { - count += EstimateTokens(messages[index].Content) + count += estimateMessageTokens(&messages[index]) } return count @@ -38,8 +44,49 @@ func EstimateInputTokens(systemPrompt string, messages []database.MessageEntity) func EstimateMessageTokens(messages []database.MessageEntity) int { tokens := 0 for index := range messages { - tokens += EstimateTokens(messages[index].Content) + tokens += estimateMessageTokens(&messages[index]) } return tokens } + +func estimateMessageTokens(message *database.MessageEntity) int { + tokens := EstimateTokens(message.Content) + for index := range message.Parts { + part := &message.Parts[index] + if part.Type == database.MessagePartText { + // Content is the text projection of multipart messages. + continue + } + + if part.Type == database.MessagePartImage { + tokens += estimateImageTokens(part.Width, part.Height) + } + } + + return tokens +} + +// Images are conservatively estimated in 512px tiles with a fixed floor. +// Malformed dimensions receive the maximum estimate rather than appearing free. +func estimateImageTokens(width, height int) int { + if width <= 0 || height <= 0 { + return maximumImageTokens + } + + tilesWide := width / imageTilePixels + if width%imageTilePixels != 0 { + tilesWide++ + } + + tilesHigh := height / imageTilePixels + if height%imageTilePixels != 0 { + tilesHigh++ + } + + if tilesWide > maximumImageTokens/imageTileTokens/tilesHigh { + return maximumImageTokens + } + + return min(maximumImageTokens, max(minimumImageTokens, tilesWide*tilesHigh*imageTileTokens)) +} diff --git a/internal/contextwindow/usage_internal_test.go b/internal/contextwindow/usage_internal_test.go index 961ac331..04ba8171 100644 --- a/internal/contextwindow/usage_internal_test.go +++ b/internal/contextwindow/usage_internal_test.go @@ -166,6 +166,24 @@ func TestMergeUsageClonesReportedBreakdownAndContributors(t *testing.T) { assert.Equal(t, "message 1", merged.TopContributors[0].Label) } +func TestEstimateMessageTokensCountsImagesConservatively(t *testing.T) { + t.Parallel() + + messages := []database.MessageEntity{{ + Timestamp: time.Time{}, Role: database.RoleUser, Content: "text", Provider: "", Model: "", + Parts: []database.MessagePartEntity{ + {Text: "text", MIMEType: "", Name: "", Type: database.MessagePartText, + Data: nil, Width: 0, Height: 0}, + {Text: "", MIMEType: "image/png", Name: "", Type: database.MessagePartImage, + Data: nil, Width: 1024, Height: 1024}, + }, + }} + + assert.Equal(t, 801, EstimateMessageTokens(messages)) + messages[0].Parts[1].Width = 0 + assert.Equal(t, 16_001, EstimateMessageTokens(messages)) +} + func TestEstimateMessageTokens(t *testing.T) { t.Parallel() @@ -184,7 +202,7 @@ func testMessageEntity(role database.Role, content string) database.MessageEntit Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, } } diff --git a/internal/contextwindow/usage_led_internal_test.go b/internal/contextwindow/usage_led_internal_test.go index 3c41affd..85f37711 100644 --- a/internal/contextwindow/usage_led_internal_test.go +++ b/internal/contextwindow/usage_led_internal_test.go @@ -115,7 +115,7 @@ func newUsageLedTestMessage(role database.Role, content string) database.Message Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, } } diff --git a/internal/database/entity.go b/internal/database/entity.go index 68d403ef..70d20398 100644 --- a/internal/database/entity.go +++ b/internal/database/entity.go @@ -48,26 +48,49 @@ const ( RoleCompactionSummary Role = "compactionSummary" ) +// MessagePartType identifies a provider-neutral message part. +type MessagePartType string + +const ( + // MessagePartText stores a textual message part. + MessagePartText MessagePartType = "text" + // MessagePartImage stores an image and its metadata. + MessagePartImage MessagePartType = "image" +) + +// MessagePartEntity is one ordered part of a durable message. +type MessagePartEntity struct { + Text string `json:"text,omitempty"` + MIMEType string `json:"mime_type,omitempty"` + Name string `json:"name,omitempty"` + Type MessagePartType `json:"type"` + Data []byte `json:"data,omitempty"` + Width int `json:"width,omitempty"` + Height int `json:"height,omitempty"` +} + // MessageEntity is the context-facing representation of an assistant message. type MessageEntity struct { - Timestamp time.Time `json:"timestamp"` - Role Role `json:"role"` - Content string `json:"content"` - Provider string `json:"provider,omitempty"` - Model string `json:"model,omitempty"` + Timestamp time.Time `json:"timestamp"` + Role Role `json:"role"` + Content string `json:"content"` + Provider string `json:"provider,omitempty"` + Model string `json:"model,omitempty"` + Parts []MessagePartEntity `json:"parts,omitempty"` } // SessionMessageEntity is the normalized durable message related to a session and entry. type SessionMessageEntity struct { - CreatedAt time.Time `json:"created_at"` - ID string `json:"id"` - SessionID string `json:"session_id"` - EntryID string `json:"entry_id"` - Sender string `json:"sender"` - Role Role `json:"role"` - Content string `json:"content"` - Provider string `json:"provider,omitempty"` - Model string `json:"model,omitempty"` + CreatedAt time.Time `json:"created_at"` + ID string `json:"id"` + SessionID string `json:"session_id"` + EntryID string `json:"entry_id"` + Sender string `json:"sender"` + Role Role `json:"role"` + Content string `json:"content"` + Provider string `json:"provider,omitempty"` + Model string `json:"model,omitempty"` + Parts []MessagePartEntity `json:"parts,omitempty"` } // SessionEntity is a persisted conversation root. @@ -84,18 +107,18 @@ type SessionEntity struct { type EntryEntity struct { CreatedAt time.Time `json:"created_at"` ParentID *string `json:"parent_id,omitempty"` - Message MessageEntity `json:"message"` - Summary string `json:"summary,omitempty"` ToolStatus string `json:"tool_status,omitempty"` - Type EntryType `json:"type"` + SessionID string `json:"session_id"` + ToolArgsJSON string `json:"tool_args_json,omitempty"` CustomType string `json:"custom_type,omitempty"` DataJSON string `json:"data_json,omitempty"` ID string `json:"id"` + Summary string `json:"summary,omitempty"` ToolName string `json:"tool_name,omitempty"` - SessionID string `json:"session_id"` - ToolArgsJSON string `json:"tool_args_json,omitempty"` + Type EntryType `json:"type"` BranchFromEntryID string `json:"branch_from_entry_id,omitempty"` CompactionFirstKeptEntryID string `json:"compaction_first_kept_entry_id,omitempty"` + Message MessageEntity `json:"message"` CompactionTokensBefore int `json:"compaction_tokens_before,omitempty"` TokenEstimate int `json:"token_estimate,omitempty"` Display bool `json:"display"` diff --git a/internal/database/migrations/00013_add_session_message_parts.sql b/internal/database/migrations/00013_add_session_message_parts.sql new file mode 100644 index 00000000..7a4c2064 --- /dev/null +++ b/internal/database/migrations/00013_add_session_message_parts.sql @@ -0,0 +1,28 @@ +-- +goose Up +CREATE UNIQUE INDEX IF NOT EXISTS idx_session_entries_id_session + ON session_entries(id, session_id); + +CREATE TABLE IF NOT EXISTS session_message_parts ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + entry_id TEXT NOT NULL, + sequence INTEGER NOT NULL CHECK(sequence >= 0), + type TEXT NOT NULL CHECK(type IN ('text', 'image')), + text TEXT NOT NULL DEFAULT '', + mime_type TEXT NOT NULL DEFAULT '', + name TEXT NOT NULL DEFAULT '', + width INTEGER NOT NULL DEFAULT 0, + height INTEGER NOT NULL DEFAULT 0, + data BLOB, + FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE CASCADE, + FOREIGN KEY (entry_id, session_id) REFERENCES session_entries(id, session_id) ON DELETE CASCADE, + UNIQUE (entry_id, sequence) +); + +CREATE INDEX IF NOT EXISTS idx_session_message_parts_session_entry_sequence + ON session_message_parts(session_id, entry_id, sequence); + +-- +goose Down +DROP INDEX IF EXISTS idx_session_message_parts_session_entry_sequence; +DROP TABLE IF EXISTS session_message_parts; +DROP INDEX IF EXISTS idx_session_entries_id_session; diff --git a/internal/database/migrations_test.go b/internal/database/migrations_test.go index f3906494..685b088d 100644 --- a/internal/database/migrations_test.go +++ b/internal/database/migrations_test.go @@ -39,6 +39,49 @@ DROP TABLE IF EXISTS workflow_agent_tasks; DROP TABLE IF EXISTS workflow_runs; ` +func TestMessagePartsMigrationUpDownAndOldSchemaUpgrade(t *testing.T) { + t.Parallel() + + connection := newMigratedThroughVersion(t, 12) + ctx := context.Background() + assertSchemaObjectCount(ctx, t, connection, 0, + `SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'session_message_parts'`) + + migrationRoot, err := database.MigrationFS() + require.NoError(t, err) + provider, err := database.NewMigrationProvider(connection, migrationRoot) + require.NoError(t, err) + _, err = provider.Up(ctx) + require.NoError(t, err) + + assertSchemaObjectExists(ctx, t, connection, + `SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'session_message_parts'`) + + const partsIndexQuery = `SELECT COUNT(*) FROM sqlite_master +WHERE type = 'index' AND name = 'idx_session_message_parts_session_entry_sequence'` + assertSchemaObjectExists(ctx, t, connection, partsIndexQuery) + + _, err = provider.Down(ctx) + require.NoError(t, err) + assertSchemaObjectCount(ctx, t, connection, 0, + `SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'session_message_parts'`) +} + +func assertSchemaObjectCount( + ctx context.Context, + t *testing.T, + connection *sql.DB, + expected int, + query string, + args ...any, +) { + t.Helper() + + var count int + require.NoError(t, connection.QueryRowContext(ctx, query, args...).Scan(&count)) + assert.Equal(t, expected, count) +} + func TestMigrateAddsCompactionOperationIdentityAfterPreviouslyDeployedVersionEleven(t *testing.T) { t.Parallel() diff --git a/internal/database/repository_helpers_internal_test.go b/internal/database/repository_helpers_internal_test.go index 59a70c44..356bc5fd 100644 --- a/internal/database/repository_helpers_internal_test.go +++ b/internal/database/repository_helpers_internal_test.go @@ -64,7 +64,7 @@ func validConstraintTestCompactionEntry(t *testing.T) *EntryEntity { Role: RoleAssistant, Content: "", Provider: "", - Model: "", + Model: "", Parts: nil, }, Summary: "summary", ToolStatus: "", diff --git a/internal/database/session_compaction_input_internal_test.go b/internal/database/session_compaction_input_internal_test.go index 6655f1c5..3a388743 100644 --- a/internal/database/session_compaction_input_internal_test.go +++ b/internal/database/session_compaction_input_internal_test.go @@ -63,7 +63,7 @@ func TestSessionRepositoryAppendCompactionValidation(t *testing.T) { Role: database.RoleUser, Content: compactionTestHistory, Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) diff --git a/internal/database/session_entry_repository.go b/internal/database/session_entry_repository.go index 2c4a0ce3..8cc805a7 100644 --- a/internal/database/session_entry_repository.go +++ b/internal/database/session_entry_repository.go @@ -52,7 +52,7 @@ func entryFromRow(row *entryRow) (*EntryEntity, error) { Role: Role(row.Role), Content: row.Content, Provider: row.Provider, - Model: row.Model, + Model: row.Model, Parts: nil, }, Summary: row.Summary, ToolStatus: row.ToolStatus, diff --git a/internal/database/session_message_parts_test.go b/internal/database/session_message_parts_test.go new file mode 100644 index 00000000..4e287193 --- /dev/null +++ b/internal/database/session_message_parts_test.go @@ -0,0 +1,251 @@ +package database_test + +import ( + "context" + "database/sql" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/omarluq/librecode/internal/database" +) + +const testImageMIME = "image/png" + +func TestSessionRepository_RoundTripsOrderedMultipartMessages(t *testing.T) { + t.Parallel() + + repository := newTestSessionRepository(t) + ctx := context.Background() + session, err := repository.CreateSession(ctx, "/work", "multipart", "") + require.NoError(t, err) + + originalData := []byte{1, 2, 3} + parts := []database.MessagePartEntity{ + {Data: nil, Text: "compare these", MIMEType: "", Name: "", Type: database.MessagePartText, Width: 0, Height: 0}, + { + Data: originalData, Text: "", MIMEType: testImageMIME, Name: "first.png", + Type: database.MessagePartImage, Width: 10, Height: 20, + }, + { + Data: []byte{4, 5}, Text: "", MIMEType: "image/jpeg", Name: "second.jpg", + Type: database.MessagePartImage, Width: 30, Height: 40, + }, + } + entry, err := repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "compare these", + Provider: "", Model: "", Parts: parts, + }) + require.NoError(t, err) + + // The repository owns its bytes after append. + originalData[0] = 99 + parts[1].Data[1] = 99 + + messages, err := repository.Messages(ctx, session.ID) + require.NoError(t, err) + require.Len(t, messages, 1) + assertMultipartParts(t, messages[0].Parts) + + transcript, err := repository.TranscriptMessages(ctx, session.ID) + require.NoError(t, err) + require.Len(t, transcript, 1) + assertMultipartParts(t, transcript[0].Parts) + + message, found, err := repository.MessageForEntry(ctx, session.ID, entry.ID) + require.NoError(t, err) + require.True(t, found) + assertMultipartParts(t, message.Parts) + + branch, err := repository.Branch(ctx, session.ID, entry.ID) + require.NoError(t, err) + require.Len(t, branch, 1) + assertMultipartParts(t, branch[0].Message.Parts) + + contextEntity, err := repository.BuildContext(ctx, session.ID, entry.ID) + require.NoError(t, err) + require.Len(t, contextEntity.Messages, 1) + assertMultipartParts(t, contextEntity.Messages[0].Parts) + + // Mutating one read cannot mutate durable data returned by a later read. + messages[0].Parts[1].Data[0] = 88 + again, err := repository.Messages(ctx, session.ID) + require.NoError(t, err) + assert.Equal(t, []byte{1, 2, 3}, again[0].Parts[1].Data) +} + +func TestSessionRepository_PersistsImageOnlyAndSupportsLegacyText(t *testing.T) { + t.Parallel() + + repository := newTestSessionRepository(t) + ctx := context.Background() + session, err := repository.CreateSession(ctx, "/work", "compatibility", "") + require.NoError(t, err) + + legacy, err := repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "legacy text", + Provider: "", Model: "", Parts: nil, + }) + require.NoError(t, err) + imageOnly, err := repository.AppendMessage(ctx, session.ID, &legacy.ID, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{{ + Data: []byte{7}, Text: "", MIMEType: testImageMIME, Name: "", + Type: database.MessagePartImage, Width: 1, Height: 1, + }}, + }) + require.NoError(t, err) + + messages, err := repository.Messages(ctx, session.ID) + require.NoError(t, err) + require.Len(t, messages, 2) + + legacyParts := []database.MessagePartEntity{{ + Data: nil, Text: "legacy text", MIMEType: "", Name: "", + Type: database.MessagePartText, Width: 0, Height: 0, + }} + assert.Equal(t, legacyParts, messages[0].Parts) + assert.Empty(t, messages[1].Content) + require.Len(t, messages[1].Parts, 1) + assert.Equal(t, database.MessagePartImage, messages[1].Parts[0].Type) + + contextEntity, err := repository.BuildContext(ctx, session.ID, imageOnly.ID) + require.NoError(t, err) + require.Len(t, contextEntity.Messages, 2) + assert.Empty(t, contextEntity.Messages[1].Content) + assert.Equal(t, []byte{7}, contextEntity.Messages[1].Parts[0].Data) +} + +func TestSessionRepository_RejectsMultipartResourceLimitBypasses(t *testing.T) { + t.Parallel() + + repository := newTestSessionRepository(t) + ctx := context.Background() + session, err := repository.CreateSession(ctx, "/work", "limits", "") + require.NoError(t, err) + + images := make([]database.MessagePartEntity, 5) + for index := range images { + images[index] = database.MessagePartEntity{ + Text: "", MIMEType: testImageMIME, Name: "", Type: database.MessagePartImage, + Data: []byte{1}, Width: 1, Height: 1, + } + } + + _, err = repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "", Provider: "", Model: "", Parts: images, + }) + require.ErrorContains(t, err, "maximum is 4") + + _, err = repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{{ + Text: "", MIMEType: testImageMIME, Name: "", Type: database.MessagePartImage, + Data: []byte{1}, Width: 40_000_001, Height: 1, + }}, + }) + require.ErrorContains(t, err, "40 megapixels") + + entries, listErr := repository.Entries(ctx, session.ID) + require.NoError(t, listErr) + assert.Empty(t, entries) +} + +func TestSessionRepository_PartInsertFailureRollsBackEntryAndMessage(t *testing.T) { + t.Parallel() + + connection, err := sql.Open(sqliteDriver(), ":memory:") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, connection.Close()) }) + connection.SetMaxOpenConns(1) + + ctx := context.Background() + require.NoError(t, database.Migrate(ctx, connection)) + require.NoError(t, database.ConfigureSQLite(ctx, connection, database.SQLiteOptions{BusyTimeout: 0})) + _, err = connection.ExecContext(ctx, `CREATE TRIGGER reject_message_part +BEFORE INSERT ON session_message_parts +BEGIN + SELECT RAISE(ABORT, 'reject message part'); +END`) + require.NoError(t, err) + + repository := database.NewSessionRepository(connection) + session, err := repository.CreateSession(ctx, "/work", "rollback", "") + require.NoError(t, err) + + _, err = repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{{ + Data: []byte{1}, Text: "", MIMEType: testImageMIME, Name: "test.png", + Type: database.MessagePartImage, Width: 1, Height: 1, + }}, + }) + require.ErrorContains(t, err, "reject message part") + + entries, err := repository.Entries(ctx, session.ID) + require.NoError(t, err) + assert.Empty(t, entries) + + messages, err := repository.Messages(ctx, session.ID) + require.NoError(t, err) + assert.Empty(t, messages) +} + +func TestSessionRepository_MessagePartsCascadeWithEntryAndSession(t *testing.T) { + t.Parallel() + + connection, err := sql.Open(sqliteDriver(), ":memory:") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, connection.Close()) }) + connection.SetMaxOpenConns(1) + + ctx := context.Background() + require.NoError(t, database.Migrate(ctx, connection)) + require.NoError(t, database.ConfigureSQLite(ctx, connection, database.SQLiteOptions{BusyTimeout: 0})) + repository := database.NewSessionRepository(connection) + + appendImage := func(name string) (*database.SessionEntity, *database.EntryEntity) { + session, createErr := repository.CreateSession(ctx, "/work", name, "") + require.NoError(t, createErr) + + entry, appendErr := repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{{ + Data: []byte{1}, Text: "", MIMEType: testImageMIME, Name: "", + Type: database.MessagePartImage, Width: 1, Height: 1, + }}, + }) + require.NoError(t, appendErr) + + return session, entry + } + countParts := func() int { + var count int + require.NoError(t, connection.QueryRowContext(ctx, `SELECT COUNT(*) FROM session_message_parts`).Scan(&count)) + + return count + } + + firstSession, firstEntry := appendImage("entry cascade") + secondSession, _ := appendImage("session cascade") + + require.Equal(t, 2, countParts()) + require.NoError(t, repository.DeleteEntryBranch(ctx, firstSession.ID, firstEntry.ID)) + assert.Equal(t, 1, countParts()) + require.NoError(t, repository.DeleteSession(ctx, secondSession.ID)) + assert.Zero(t, countParts()) +} + +func assertMultipartParts(t *testing.T, parts []database.MessagePartEntity) { + t.Helper() + require.Len(t, parts, 3) + assert.Equal(t, database.MessagePartText, parts[0].Type) + assert.Equal(t, "compare these", parts[0].Text) + assert.Equal(t, database.MessagePartImage, parts[1].Type) + assert.Equal(t, []byte{1, 2, 3}, parts[1].Data) + assert.Equal(t, "first.png", parts[1].Name) + assert.Equal(t, database.MessagePartImage, parts[2].Type) + assert.Equal(t, []byte{4, 5}, parts[2].Data) +} diff --git a/internal/database/session_message_repository.go b/internal/database/session_message_repository.go index c7324a6a..347762bf 100644 --- a/internal/database/session_message_repository.go +++ b/internal/database/session_message_repository.go @@ -3,11 +3,24 @@ package database import ( "context" "errors" + "fmt" + "strings" "github.com/samber/oops" "github.com/vingarcia/ksql" ) +type sessionMessagePartRow struct { + EntryID string `ksql:"entry_id"` + Type string `ksql:"type"` + Text string `ksql:"text"` + MIMEType string `ksql:"mime_type"` + Name string `ksql:"name"` + Data []byte `ksql:"data"` + Width int `ksql:"width"` + Height int `ksql:"height"` +} + type sessionMessageRow struct { ID string `ksql:"id"` SessionID string `ksql:"session_id"` @@ -35,7 +48,7 @@ func sessionMessageFromRow(row *sessionMessageRow) (*SessionMessageEntity, error Role: Role(row.Role), Content: row.Content, Provider: row.Provider, - Model: row.Model, + Model: row.Model, Parts: nil, }, nil } @@ -51,17 +64,7 @@ FROM session_messages WHERE session_id = ? ORDER BY created_at ASC` - rows := []sessionMessageRow{} - if err := repository.sql.Query(ctx, &rows, query, sessionID); err != nil { - return nil, oops.In("database").Code("list_messages").Wrapf(err, "query session messages") - } - - messages, err := sessionMessagesFromRows(rows) - if err != nil { - return nil, oops.In("database").Code("scan_message").Wrapf(err, "scan session messages") - } - - return messages, nil + return repository.querySessionMessages(ctx, sessionID, query, "messages") } // TranscriptMessages returns displayable normalized messages for a session in creation order. @@ -76,20 +79,29 @@ JOIN session_entries AS e ON e.id = m.entry_id AND e.session_id = m.session_id WHERE m.session_id = ? AND e.display = 1 ORDER BY m.created_at ASC` + return repository.querySessionMessages(ctx, sessionID, query, "transcript_messages") +} + +func (repository *SessionRepository) querySessionMessages( + ctx context.Context, + sessionID string, + query string, + operation string, +) ([]SessionMessageEntity, error) { + operationLabel := strings.ReplaceAll(operation, "_", " ") + rows := []sessionMessageRow{} if err := repository.sql.Query(ctx, &rows, query, sessionID); err != nil { - return nil, oops.In("database").Code("list_transcript_messages").Wrapf( - err, - "query transcript messages", - ) + return nil, oops.In("database").Code("list_"+operation).Wrapf(err, "query %s", operationLabel) } messages, err := sessionMessagesFromRows(rows) if err != nil { - return nil, oops.In("database").Code("scan_transcript_message").Wrapf( - err, - "scan transcript messages", - ) + return nil, oops.In("database").Code("scan_"+operation).Wrapf(err, "scan %s", operationLabel) + } + + if err := repository.hydrateSessionMessages(ctx, sessionID, messages); err != nil { + return nil, err } return messages, nil @@ -120,7 +132,12 @@ WHERE session_id = ? AND entry_id = ?` return nil, false, oops.In("database").Code("scan_message").Wrapf(err, "scan session message") } - return message, true, nil + messages := []SessionMessageEntity{*message} + if err := repository.hydrateSessionMessages(ctx, sessionID, messages); err != nil { + return nil, false, err + } + + return &messages[0], true, nil } func (repository *SessionRepository) appendEntryMessage( @@ -158,6 +175,21 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)` return oops.In("database").Code("append_message").Wrapf(err, "append session message") } + const insertPart = ` +INSERT INTO session_message_parts + (id, session_id, entry_id, sequence, type, text, mime_type, name, width, height, data) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)` + + for sequence := range message.Parts { + part := &message.Parts[sequence] + if _, err := transaction.Exec(ctx, insertPart, + newEntryID(), message.SessionID, message.EntryID, sequence, string(part.Type), + part.Text, part.MIMEType, part.Name, part.Width, part.Height, cloneBytes(part.Data), + ); err != nil { + return oops.In("database").Code("append_message_part").Wrapf(err, "append session message part") + } + } + return nil } @@ -181,7 +213,143 @@ func sessionMessageFromEntry(entry *EntryEntity) SessionMessageEntity { Content: entry.Message.Content, Provider: entry.Message.Provider, Model: entry.Message.Model, + Parts: cloneMessageParts(entry.Message.Parts), + } +} + +func (repository *SessionRepository) hydrateSessionMessages( + ctx context.Context, + sessionID string, + messages []SessionMessageEntity, +) error { + entryIDs := make([]string, len(messages)) + for index := range messages { + entryIDs[index] = messages[index].EntryID + } + + partsByEntry, err := repository.messagePartsForEntries(ctx, sessionID, entryIDs) + if err != nil { + return err + } + + for index := range messages { + messages[index].Parts = partsOrLegacyText(partsByEntry[messages[index].EntryID], messages[index].Content) + } + + return nil +} + +func (repository *SessionRepository) hydrateEntryMessages( + ctx context.Context, + sessionID string, + entries []EntryEntity, +) error { + entryIDs := make([]string, 0, len(entries)) + for index := range entries { + if entryCarriesMessage(&entries[index]) { + entryIDs = append(entryIDs, entries[index].ID) + } } + + partsByEntry, err := repository.messagePartsForEntries(ctx, sessionID, entryIDs) + if err != nil { + return err + } + + for index := range entries { + if entryCarriesMessage(&entries[index]) { + parts := partsByEntry[entries[index].ID] + entries[index].Message.Parts = partsOrLegacyText(parts, entries[index].Message.Content) + } + } + + return nil +} + +func (repository *SessionRepository) messagePartsForEntries( + ctx context.Context, + sessionID string, + entryIDs []string, +) (map[string][]MessagePartEntity, error) { + const entryIDBatchSize = 900 + + partsByEntry := make(map[string][]MessagePartEntity, len(entryIDs)) + for start := 0; start < len(entryIDs); start += entryIDBatchSize { + end := min(start+entryIDBatchSize, len(entryIDs)) + if err := repository.appendMessagePartsBatch(ctx, sessionID, entryIDs[start:end], partsByEntry); err != nil { + return nil, err + } + } + + return partsByEntry, nil +} + +func (repository *SessionRepository) appendMessagePartsBatch( + ctx context.Context, + sessionID string, + entryIDs []string, + partsByEntry map[string][]MessagePartEntity, +) error { + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(entryIDs)), ",") + query := fmt.Sprintf(` +SELECT entry_id, type, text, data, mime_type, name, width, height +FROM session_message_parts +WHERE session_id = ? AND entry_id IN (%s) +ORDER BY entry_id, sequence`, placeholders) + args := make([]any, 0, len(entryIDs)+1) + + args = append(args, sessionID) + for _, entryID := range entryIDs { + args = append(args, entryID) + } + + rows := []sessionMessagePartRow{} + if err := repository.sql.Query(ctx, &rows, query, args...); err != nil { + return oops.In("database").Code("list_message_parts").Wrapf(err, "query message parts") + } + + for index := range rows { + row := &rows[index] + partsByEntry[row.EntryID] = append(partsByEntry[row.EntryID], MessagePartEntity{ + Type: MessagePartType(row.Type), Text: row.Text, Data: cloneBytes(row.Data), + MIMEType: row.MIMEType, Name: row.Name, Width: row.Width, Height: row.Height, + }) + } + + return nil +} + +func partsOrLegacyText(parts []MessagePartEntity, content string) []MessagePartEntity { + if len(parts) == 0 && strings.TrimSpace(content) != "" { + return []MessagePartEntity{{ + Data: nil, Text: content, MIMEType: "", Name: "", Type: MessagePartText, Width: 0, Height: 0, + }} + } + + return cloneMessageParts(parts) +} + +func cloneMessageParts(parts []MessagePartEntity) []MessagePartEntity { + if parts == nil { + return nil + } + + cloned := make([]MessagePartEntity, len(parts)) + copy(cloned, parts) + + for index := range cloned { + cloned[index].Data = cloneBytes(parts[index].Data) + } + + return cloned +} + +func cloneBytes(data []byte) []byte { + if data == nil { + return nil + } + + return append([]byte(nil), data...) } func senderIdentity(entry *EntryEntity) string { diff --git a/internal/database/session_repository_test.go b/internal/database/session_repository_test.go index fcd7deac..6331582c 100644 --- a/internal/database/session_repository_test.go +++ b/internal/database/session_repository_test.go @@ -39,7 +39,7 @@ func TestSessionRepository_AppendsMessagesInSessionTree(t *testing.T) { Role: database.RoleUser, Content: testHello, Provider: "", - Model: "", + Model: "", Parts: nil, } firstEntry, err := repository.AppendMessage(ctx, createdSession.ID, nil, &firstMessage) require.NoError(t, err) @@ -50,7 +50,7 @@ func TestSessionRepository_AppendsMessagesInSessionTree(t *testing.T) { Role: database.RoleAssistant, Content: "hi", Provider: "local", - Model: "librecode", + Model: "librecode", Parts: nil, } secondEntry, err := repository.AppendMessage(ctx, createdSession.ID, &firstEntry.ID, &secondMessage) require.NoError(t, err) @@ -115,7 +115,7 @@ func TestSessionRepository_TranscriptMessagesHonorsDisplayMetadata(t *testing.T) Role: database.RoleUser, Content: testVisible, Provider: "", - Model: "", + Model: "", Parts: nil, }, &modelFacing, &visible) require.NoError(t, err) second, err := repository.AppendMessageWithDisplay(ctx, session.ID, &first.ID, &database.MessageEntity{ @@ -123,7 +123,7 @@ func TestSessionRepository_TranscriptMessagesHonorsDisplayMetadata(t *testing.T) Role: database.RoleAssistant, Content: "hidden", Provider: "", - Model: "", + Model: "", Parts: nil, }, &modelFacing, &hidden) require.NoError(t, err) _, err = repository.AppendMessage(ctx, session.ID, &second.ID, &database.MessageEntity{ @@ -131,7 +131,7 @@ func TestSessionRepository_TranscriptMessagesHonorsDisplayMetadata(t *testing.T) Role: database.RoleAssistant, Content: "default visible", Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) @@ -163,7 +163,7 @@ func TestSessionRepository_AppendMessageWithDisplayDefaultsIndependently(t *test hidden := false entry, err := repository.AppendMessageWithDisplay(ctx, session.ID, nil, &database.MessageEntity{ Timestamp: time.Time{}, Role: database.RoleUser, Content: "hidden but model-facing by role default", - Provider: "", Model: "", + Provider: "", Model: "", Parts: nil, }, nil, &hidden) require.NoError(t, err) assert.False(t, entry.Display) @@ -172,7 +172,7 @@ func TestSessionRepository_AppendMessageWithDisplayDefaultsIndependently(t *test visible := true entry, err = repository.AppendMessageWithDisplay(ctx, session.ID, &entry.ID, &database.MessageEntity{ Timestamp: time.Time{}, Role: database.RoleAssistant, Content: "shown", - Provider: "", Model: "", + Provider: "", Model: "", Parts: nil, }, nil, &visible) require.NoError(t, err) assert.True(t, entry.Display) @@ -193,7 +193,7 @@ func TestSessionRepository_TranscriptMessagesWrapsMalformedRows(t *testing.T) { require.NoError(t, err) _, err = repository.AppendMessage(context.Background(), session.ID, nil, &database.MessageEntity{ Timestamp: time.Time{}, Role: database.RoleUser, Content: "malformed", - Provider: "", Model: "", + Provider: "", Model: "", Parts: nil, }) require.NoError(t, err) _, err = connection.ExecContext(context.Background(), `UPDATE session_messages SET created_at = 'invalid'`) @@ -248,7 +248,7 @@ func TestSessionRepository_AppendMessagePreservesInputTimestamp(t *testing.T) { Role: database.RoleUser, Content: testHello, Provider: "", - Model: "", + Model: "", Parts: nil, } expectedTimestamp := testCase.timestamp.UTC() diff --git a/internal/database/session_store.go b/internal/database/session_store.go index 1367d839..82120a5c 100644 --- a/internal/database/session_store.go +++ b/internal/database/session_store.go @@ -79,6 +79,10 @@ func (repository *SessionRepository) Branch(ctx context.Context, sessionID, entr Errorf("entry %q not found in session %q", entryID, sessionID) } + if err := repository.hydrateEntryMessages(ctx, sessionID, entries); err != nil { + return nil, err + } + return entries, nil } @@ -251,6 +255,7 @@ type appendEntryOptions struct { provider string summary string role Role + parts []MessagePartEntity } func newAppendEntryOptions() *appendEntryOptions { @@ -297,6 +302,7 @@ func (repository *SessionRepository) entryFromAppendOptions( Content: options.content, Provider: options.provider, Model: options.model, + Parts: cloneMessageParts(options.parts), }) entry.CustomType = options.customType @@ -375,12 +381,7 @@ func (repository *SessionRepository) AppendMessageWithDisplay( modelFacing *bool, display *bool, ) (*EntryEntity, error) { - options := newAppendEntryOptions() - options.content = message.Content - options.model = message.Model - options.provider = message.Provider - options.role = message.Role - options.timestamp = message.Timestamp + options := appendOptionsFromMessage(message) options.modelFacing = modelFacing options.display = display @@ -396,16 +397,23 @@ func (repository *SessionRepository) AppendMessageWithMetadata( modelFacing *bool, usage *EntryTokenUsageEntity, ) (*EntryEntity, error) { + options := appendOptionsFromMessage(message) + options.modelFacing = modelFacing + options.usage = usage + + return repository.appendBuiltEntry(ctx, sessionID, parentID, EntryTypeMessage, options) +} + +func appendOptionsFromMessage(message *MessageEntity) *appendEntryOptions { options := newAppendEntryOptions() options.content = message.Content options.model = message.Model options.provider = message.Provider options.role = message.Role options.timestamp = message.Timestamp - options.modelFacing = modelFacing - options.usage = usage + options.parts = message.Parts - return repository.appendBuiltEntry(ctx, sessionID, parentID, EntryTypeMessage, options) + return options } // AppendCustom appends extension state that does not participate in prompt context. @@ -711,7 +719,7 @@ func appendBranchSummaryContext(contextEntity *SessionContextEntity, entry *Entr Role: RoleBranchSummary, Content: entry.Summary, Provider: "", - Model: "", + Model: "", Parts: nil, }) return nil @@ -757,6 +765,7 @@ func compactionSummaryMessages(entry *EntryEntity) []MessageEntity { Content: entry.Summary, Provider: "", Model: "", + Parts: nil, }} } diff --git a/internal/database/session_usage_test.go b/internal/database/session_usage_test.go index 173af359..ebf77296 100644 --- a/internal/database/session_usage_test.go +++ b/internal/database/session_usage_test.go @@ -103,6 +103,6 @@ func newUsageTestMessage(role database.Role, content, provider, model string) *d Role: role, Content: content, Provider: provider, - Model: model, + Model: model, Parts: nil, } } diff --git a/internal/database/sqlite_contention_internal_test.go b/internal/database/sqlite_contention_internal_test.go index d2c4820d..cef6039d 100644 --- a/internal/database/sqlite_contention_internal_test.go +++ b/internal/database/sqlite_contention_internal_test.go @@ -50,7 +50,7 @@ func TestSessionRepositoryConcurrentWritersWaitForBusyDatabase(t *testing.T) { Role: database.RoleUser, Content: strings.Repeat("x", writerIndex+entryIndex+1), Provider: "", - Model: "", + Model: "", Parts: nil, }) appendErrors <- appendErr } @@ -89,7 +89,7 @@ func TestSessionRepositoryConcurrentCompactionsChooseOneWinner(t *testing.T) { Role: database.RoleUser, Content: compactionTestHistory, Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) @@ -162,7 +162,7 @@ func TestSessionRepositoryCompactionOperationIsIdempotent(t *testing.T) { Role: database.RoleUser, Content: compactionTestHistory, Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) @@ -183,7 +183,7 @@ func TestSessionRepositoryCompactionOperationIsIdempotent(t *testing.T) { Role: database.RoleUser, Content: "different branch target", Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) mismatched, err := secondary.AppendCompaction( diff --git a/internal/database/task_validation_internal_test.go b/internal/database/task_validation_internal_test.go index c8dcfc50..f1d40a38 100644 --- a/internal/database/task_validation_internal_test.go +++ b/internal/database/task_validation_internal_test.go @@ -332,7 +332,7 @@ func validSessionEntity(entityID string, now time.Time) SessionEntity { func validEntryEntity(entityID string, now time.Time) EntryEntity { return EntryEntity{CreatedAt: now, ParentID: nil, Message: MessageEntity{Timestamp: time.Time{}, Role: "", - Content: "", Provider: "", Model: ""}, Summary: "", ToolStatus: "", + Content: "", Provider: "", Model: "", Parts: nil}, Summary: "", ToolStatus: "", Type: EntryTypeMessage, CustomType: "", DataJSON: `{}`, ID: entityID, ToolName: "", SessionID: entityID, ToolArgsJSON: "", BranchFromEntryID: "", CompactionFirstKeptEntryID: "", CompactionTokensBefore: 0, TokenEstimate: 0, Display: false, ModelFacing: false} @@ -340,7 +340,7 @@ func validEntryEntity(entityID string, now time.Time) EntryEntity { func validMessageEntity(entityID string, now time.Time) SessionMessageEntity { return SessionMessageEntity{CreatedAt: now, ID: entityID, SessionID: entityID, EntryID: entityID, - Sender: string(RoleUser), Role: RoleUser, Content: "", Provider: "", Model: ""} + Sender: string(RoleUser), Role: RoleUser, Content: "", Provider: "", Model: "", Parts: nil} } func validTaskEntity(entityID string, now time.Time) TaskEntity { diff --git a/internal/database/test_helpers_test.go b/internal/database/test_helpers_test.go index a5c93d10..829e3d18 100644 --- a/internal/database/test_helpers_test.go +++ b/internal/database/test_helpers_test.go @@ -109,7 +109,7 @@ func (h sessionTestHelper) appendMessageAt( Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(h.t, err) diff --git a/internal/database/validation.go b/internal/database/validation.go index 514e9733..be501310 100644 --- a/internal/database/validation.go +++ b/internal/database/validation.go @@ -4,6 +4,7 @@ import ( "encoding/json" "errors" "fmt" + "mime" "strings" "time" @@ -57,6 +58,13 @@ func validateEntryEntity(entity *EntryEntity) error { return nil } +const ( + maxMessageImages = 4 + maxMessageImageBytes = 5 << 20 + maxMessageImageTotal = 20 << 20 + maxMessageImagePixels = 40_000_000 +) + func validateSessionMessageEntity(entity *SessionMessageEntity) error { if err := validateUUIDv7("message.id", entity.ID); err != nil { return err @@ -78,9 +86,110 @@ func validateSessionMessageEntity(entity *SessionMessageEntity) error { return err } + if err := validateMessageParts(entity.Parts); err != nil { + return err + } + + if len(entity.Parts) > 0 && entity.Content != messagePartsText(entity.Parts) { + return errors.New("message.content must match the text projection of message.parts") + } + return validateRequiredTime("message.created_at", entity.CreatedAt) } +func messagePartsText(parts []MessagePartEntity) string { + var text strings.Builder + + for index := range parts { + if parts[index].Type == MessagePartText { + text.WriteString(parts[index].Text) + } + } + + return text.String() +} + +func validateMessageParts(parts []MessagePartEntity) error { + imageCount := 0 + imageBytes := 0 + + for index := range parts { + part := &parts[index] + if err := validateMessagePartEntity(part); err != nil { + return fmt.Errorf("message.parts[%d]: %w", index, err) + } + + if part.Type == MessagePartImage { + imageCount++ + imageBytes += len(part.Data) + } + } + + if imageCount > maxMessageImages { + return fmt.Errorf("message has %d images; maximum is %d", imageCount, maxMessageImages) + } + + if imageBytes > maxMessageImageTotal { + return errors.New("message image data exceeds the 20 MiB limit") + } + + return nil +} + +func validateMessagePartEntity(part *MessagePartEntity) error { + switch part.Type { + case MessagePartText: + return validateTextMessagePart(part) + case MessagePartImage: + return validateImageMessagePart(part) + default: + return fmt.Errorf("unsupported message part type %q", part.Type) + } +} + +func validateTextMessagePart(part *MessagePartEntity) error { + if strings.TrimSpace(part.Text) == "" { + return errors.New("text part must have text") + } + + if len(part.Data) != 0 { + return errors.New("text part must not have binary data") + } + + return nil +} + +func validateImageMessagePart(part *MessagePartEntity) error { + if len(part.Data) == 0 { + return errors.New("image part must have binary data") + } + + if part.Text != "" { + return errors.New("image part must not have text") + } + + if !validImageMIMEType(part.MIMEType) { + return errors.New("image part must have a normalized image MIME type") + } + + if len(part.Data) > maxMessageImageBytes { + return errors.New("image part exceeds the 5 MiB limit") + } + + if part.Width <= 0 || part.Height <= 0 || part.Width > maxMessageImagePixels/part.Height { + return errors.New("image part dimensions must be positive and at most 40 megapixels") + } + + return nil +} + +func validImageMIMEType(value string) bool { + mediaType, _, err := mime.ParseMediaType(value) + + return err == nil && mediaType == value && value == strings.ToLower(value) && + strings.HasPrefix(value, "image/") && len(value) > len("image/") +} + func validateTaskEntity(entity *TaskEntity) error { if err := validateUUIDv7("task.id", entity.ID); err != nil { return err diff --git a/internal/extension/lua_values.go b/internal/extension/lua_values.go index 85339417..1c7c42e4 100644 --- a/internal/extension/lua_values.go +++ b/internal/extension/lua_values.go @@ -83,6 +83,13 @@ func newLuaValue(state *lua.LState, value any) luaValue { return luaValue{value: stringMapToLuaTable(state, typedValue)} case []any: return luaValue{value: sliceToLuaTable(state, typedValue)} + case []map[string]any: + values := make([]any, len(typedValue)) + for index := range typedValue { + values[index] = typedValue[index] + } + + return luaValue{value: sliceToLuaTable(state, values)} case []string: return luaValue{value: stringSliceToLuaTable(state, typedValue)} default: diff --git a/internal/model/message_filter.go b/internal/model/message_filter.go index 7972b8c2..7c745c5d 100644 --- a/internal/model/message_filter.go +++ b/internal/model/message_filter.go @@ -25,6 +25,6 @@ func emptyMessage() database.MessageEntity { Role: "", Content: "", Provider: "", - Model: "", + Model: "", Parts: nil, } } diff --git a/internal/model/messages.go b/internal/model/messages.go index 40095c3a..4e9dd8a7 100644 --- a/internal/model/messages.go +++ b/internal/model/messages.go @@ -27,7 +27,7 @@ func IsFacingRole(role database.Role) bool { // FacingMessage converts persisted summary roles into model-facing user messages. func FacingMessage(message *database.MessageEntity) database.MessageEntity { - converted := *message + converted := cloneMessage(message) switch message.Role { case database.RoleCompactionSummary: converted.Role = database.RoleUser @@ -49,5 +49,36 @@ func FacingMessage(message *database.MessageEntity) database.MessageEntity { // IsFacingMessage reports whether a persisted message has model-facing content. func IsFacingMessage(message *database.MessageEntity) bool { - return message != nil && IsFacingRole(message.Role) && strings.TrimSpace(message.Content) != "" + if message == nil || !IsFacingRole(message.Role) { + return false + } + + if strings.TrimSpace(message.Content) != "" { + return true + } + + for index := range message.Parts { + part := &message.Parts[index] + if part.Type == database.MessagePartImage && len(part.Data) > 0 { + return true + } + } + + return false +} + +func cloneMessage(message *database.MessageEntity) database.MessageEntity { + converted := *message + if message.Parts == nil { + return converted + } + + converted.Parts = make([]database.MessagePartEntity, len(message.Parts)) + copy(converted.Parts, message.Parts) + + for index := range converted.Parts { + converted.Parts[index].Data = append([]byte(nil), message.Parts[index].Data...) + } + + return converted } diff --git a/internal/model/messages_test.go b/internal/model/messages_test.go index b09bd9f8..f9137ffc 100644 --- a/internal/model/messages_test.go +++ b/internal/model/messages_test.go @@ -110,11 +110,32 @@ func TestIsFacingMessageHandlesNilAndBlankContent(t *testing.T) { blank := testMessage(database.RoleUser, " \n\t ") toolResult := testMessage(database.RoleToolResult, "visible text") user := testMessage(database.RoleUser, "visible text") + imageOnly := testMessage(database.RoleUser, "") + imageOnly.Parts = []database.MessagePartEntity{{ + Text: "", MIMEType: "image/png", Name: "", Type: database.MessagePartImage, + Data: []byte{1}, Width: 1, Height: 1, + }} assert.False(t, model.IsFacingMessage(nil)) assert.False(t, model.IsFacingMessage(&blank)) assert.False(t, model.IsFacingMessage(&toolResult)) assert.True(t, model.IsFacingMessage(&user)) + assert.True(t, model.IsFacingMessage(&imageOnly)) +} + +func TestFacingMessageDeepCopiesImageData(t *testing.T) { + t.Parallel() + + message := testMessage(database.RoleUser, "") + message.Parts = []database.MessagePartEntity{{ + Text: "", MIMEType: "image/png", Name: "", Type: database.MessagePartImage, + Data: []byte{1, 2}, Width: 1, Height: 1, + }} + + converted := model.FacingMessage(&message) + converted.Parts[0].Data[0] = 9 + + assert.Equal(t, byte(1), message.Parts[0].Data[0]) } func testMessage(role database.Role, content string) database.MessageEntity { @@ -123,6 +144,6 @@ func testMessage(role database.Role, content string) database.MessageEntity { Role: role, Content: content, Provider: "provider", - Model: "model", + Model: "model", Parts: nil, } } diff --git a/internal/provider/anthropic.go b/internal/provider/anthropic.go index c16f4749..a5700d86 100644 --- a/internal/provider/anthropic.go +++ b/internal/provider/anthropic.go @@ -18,8 +18,13 @@ func (client *HTTPCompletionClient) completeAnthropic( ctx context.Context, request *CompletionRequest, ) (*llm.Response, error) { + messages, err := anthropicMessages(request.Request.Messages) + if err != nil { + return nil, err + } + state := anthropicLoopState{ - messages: anthropicMessages(request.Request.Messages), + messages: messages, endpoint: joinEndpoint(request.Request.Model.BaseURL, "/v1/messages"), result: newResponse(), } @@ -436,21 +441,33 @@ func anthropicToolResultMessage(calls []ToolCall, events []ToolEvent) (map[strin return map[string]any{jsonRoleKey: jsonUserRole, jsonContentKey: blocks}, nil } -func anthropicMessages(messages []llm.Message) []map[string]any { +func anthropicMessages(messages []llm.Message) ([]map[string]any, error) { + if err := validateImageMessages(messages); err != nil { + return nil, err + } + output := []map[string]any{} for _, message := range messages { role, ok := anthropicRole(message.Role) - content := messageText(message) - if !ok || content == "" { + if !ok { + continue + } + + var content any = messageText(message) + if message.Role == llm.RoleUser { + content = anthropicUserContent(message) + } + + if content == "" { continue } output = append(output, map[string]any{jsonRoleKey: role, jsonContentKey: content}) } - return output + return output, nil } // AnthropicToolsFromDefinitions returns Anthropic tool declarations for definitions. diff --git a/internal/provider/anthropic_mapping_internal_test.go b/internal/provider/anthropic_mapping_internal_test.go index 53dd3191..950604fc 100644 --- a/internal/provider/anthropic_mapping_internal_test.go +++ b/internal/provider/anthropic_mapping_internal_test.go @@ -167,8 +167,9 @@ func TestAnthropicMessagesAndRoles(t *testing.T) { request := emptyCompletionRequest() setTestRequestMessages(request, mixedReplayMessages()) - converted := anthropicMessages(request.Request.Messages) + converted, err := anthropicMessages(request.Request.Messages) + require.NoError(t, err) assert.Len(t, converted, 7) assert.JSONEq(t, jsonString(jsonUserRole), jsonString(converted[0][jsonRoleKey])) assert.JSONEq(t, jsonString(jsonAssistantRole), jsonString(converted[1][jsonRoleKey])) @@ -178,6 +179,53 @@ func TestAnthropicMessagesAndRoles(t *testing.T) { assert.Empty(t, mapped) } +func TestAnthropicMultipartPayloadShapes(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + message llm.Message + want []map[string]any + }{ + { + name: testImageWithText, + message: testImageMessage(), + want: []map[string]any{ + {jsonTypeKey: jsonTextKey, jsonTextKey: testImagePrompt}, + {jsonTypeKey: "image", "source": map[string]any{ + jsonTypeKey: "base64", "media_type": testImageMIME, "data": testImageData, + }}, + }, + }, + { + name: testImageOnly, + message: testImageOnlyMessage(llm.RoleUser, testImageMIME, testImageData), + want: []map[string]any{{jsonTypeKey: "image", "source": map[string]any{ + jsonTypeKey: "base64", "media_type": testImageMIME, "data": testImageData, + }}}, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + messages, err := anthropicMessages([]llm.Message{ + test.message, + llm.TextMessage(llm.RoleAssistant, testAssistantReplay), + llm.TextMessage(llm.RoleTool, "tool output"), + }) + + require.NoError(t, err) + require.Len(t, messages, 2) + assert.Equal(t, test.want, messages[0][jsonContentKey]) + assert.Equal(t, map[string]any{ + jsonRoleKey: jsonAssistantRole, jsonContentKey: testAssistantReplay, + }, messages[1]) + }) + } +} + func TestAnthropicLocalToolNameFallbacks(t *testing.T) { t.Parallel() diff --git a/internal/provider/image_content.go b/internal/provider/image_content.go new file mode 100644 index 00000000..7b071727 --- /dev/null +++ b/internal/provider/image_content.go @@ -0,0 +1,159 @@ +package provider + +import ( + "encoding/base64" + "fmt" + + "github.com/samber/oops" + + "github.com/omarluq/librecode/internal/llm" +) + +func validateImageMessages(messages []llm.Message) error { + for messageIndex := range messages { + message := &messages[messageIndex] + for partIndex := range message.Content { + part := &message.Content[partIndex] + if err := validateImagePart(message.Role, part, messageIndex, partIndex); err != nil { + return err + } + } + } + + return nil +} + +func validateImagePart(role llm.Role, part *llm.Part, messageIndex, partIndex int) error { + if part.Type != llm.PartImage { + return nil + } + + if role != llm.RoleUser { + return imageConversionError(messageIndex, partIndex, "image content is only supported for user messages") + } + + if part.Data == "" { + return imageConversionError(messageIndex, partIndex, "image content requires base64 data") + } + + if !supportedImageMIMEType(part.MIMEType) { + return imageConversionError(messageIndex, partIndex, "unsupported image MIME type") + } + + if _, err := base64.StdEncoding.DecodeString(part.Data); err != nil { + return imageConversionError(messageIndex, partIndex, "image content contains malformed base64 data") + } + + return nil +} + +func supportedImageMIMEType(mimeType string) bool { + switch mimeType { + case "image/gif", "image/jpeg", "image/png", "image/webp": + return true + default: + return false + } +} + +func imageConversionError(messageIndex, partIndex int, text string) error { + return oops.In("provider").Code("invalid_image_content"). + With("message_index", messageIndex).With("part_index", partIndex).Errorf("%s", text) +} + +const ( + jsonImageURLKey = "image_url" + jsonImageURLValueKey = "url" + openAIInputImageType = "input_image" + anthropicImageType = "image" + anthropicBase64Type = "base64" + anthropicSourceKey = "source" + anthropicMediaTypeKey = "media_type" + encodedImageDataKey = "data" +) + +func openAIDataURL(part *llm.Part) string { + return fmt.Sprintf("data:%s;base64,%s", part.MIMEType, part.Data) +} + +func openAIResponseUserContent(message llm.Message) []map[string]any { + blocks := make([]map[string]any, 0, len(message.Content)) + for index := range message.Content { + part := &message.Content[index] + switch part.Type { + case llm.PartText: + if part.Text != "" { + blocks = append(blocks, map[string]any{jsonTypeKey: "input_text", jsonTextKey: part.Text}) + } + case llm.PartImage: + blocks = append(blocks, map[string]any{ + jsonTypeKey: openAIInputImageType, jsonImageURLKey: openAIDataURL(part), + }) + case llm.PartReasoning, llm.PartFile, llm.PartSource, llm.PartToolCall, llm.PartToolResult: + continue + } + } + + return blocks +} + +func openAIChatUserContent(message llm.Message) any { + hasImage := false + for index := range message.Content { + hasImage = hasImage || message.Content[index].Type == llm.PartImage + } + + if !hasImage { + return messageText(message) + } + + blocks := make([]map[string]any, 0, len(message.Content)) + for index := range message.Content { + part := &message.Content[index] + switch part.Type { + case llm.PartText: + if part.Text != "" { + blocks = append(blocks, map[string]any{jsonTypeKey: jsonTextKey, jsonTextKey: part.Text}) + } + case llm.PartImage: + blocks = append(blocks, map[string]any{ + jsonTypeKey: jsonImageURLKey, + jsonImageURLKey: map[string]any{jsonImageURLValueKey: openAIDataURL(part)}, + }) + case llm.PartReasoning, llm.PartFile, llm.PartSource, llm.PartToolCall, llm.PartToolResult: + continue + } + } + + return blocks +} + +func anthropicUserContent(message llm.Message) any { + hasImage := false + for index := range message.Content { + hasImage = hasImage || message.Content[index].Type == llm.PartImage + } + + if !hasImage { + return messageText(message) + } + + blocks := make([]map[string]any, 0, len(message.Content)) + for index := range message.Content { + part := &message.Content[index] + switch part.Type { + case llm.PartText: + if part.Text != "" { + blocks = append(blocks, map[string]any{jsonTypeKey: jsonTextKey, jsonTextKey: part.Text}) + } + case llm.PartImage: + blocks = append(blocks, map[string]any{jsonTypeKey: anthropicImageType, anthropicSourceKey: map[string]any{ + jsonTypeKey: anthropicBase64Type, anthropicMediaTypeKey: part.MIMEType, encodedImageDataKey: part.Data, + }}) + case llm.PartReasoning, llm.PartFile, llm.PartSource, llm.PartToolCall, llm.PartToolResult: + continue + } + } + + return blocks +} diff --git a/internal/provider/image_content_internal_test.go b/internal/provider/image_content_internal_test.go new file mode 100644 index 00000000..28396a51 --- /dev/null +++ b/internal/provider/image_content_internal_test.go @@ -0,0 +1,194 @@ +package provider + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/omarluq/librecode/internal/llm" +) + +const ( + testImageMIME = "image/png" + testImageData = "AQID" + testImageDataURL = "data:image/png;base64,AQID" + testImagePrompt = "describe image" + testImageWithText = "text and image" + testImageOnly = "image only" + testAssistantReplay = "answer" +) + +func testImageMessage() llm.Message { + return llm.Message{Metadata: nil, Role: llm.RoleUser, Content: []llm.Part{ + llm.TextPart(testImagePrompt), + {Metadata: nil, ToolCall: nil, ToolResult: nil, Type: llm.PartImage, + Text: "", MIMEType: testImageMIME, Data: testImageData}, + }} +} + +func testImageOnlyMessage(role llm.Role, mimeType, data string) llm.Message { + return llm.Message{Metadata: nil, Role: role, Content: []llm.Part{{ + Metadata: nil, ToolCall: nil, ToolResult: nil, Type: llm.PartImage, + Text: "", MIMEType: mimeType, Data: data, + }}} +} + +func TestMultipartProviderMappings(t *testing.T) { + t.Parallel() + + message := testImageMessage() + responses, err := openAIResponseInput([]llm.Message{message}) + require.NoError(t, err) + + responseObject, converted := responses[0].(map[string]any) + require.True(t, converted) + + responseContent, contentConverted := responseObject[jsonContentKey].([]map[string]any) + require.True(t, contentConverted) + assert.Equal(t, "input_text", responseContent[0][jsonTypeKey]) + assert.Equal(t, "input_image", responseContent[1][jsonTypeKey]) + assert.Equal(t, "data:image/png;base64,AQID", responseContent[1][jsonImageURLKey]) + + request := emptyCompletionRequest() + setTestRequestMessages(request, []llm.Message{message}) + chat, err := openAIChatMessages(request) + require.NoError(t, err) + + chatContent, chatConverted := chat[0][jsonContentKey].([]map[string]any) + require.True(t, chatConverted) + + imageURL, imageURLConverted := chatContent[1][jsonImageURLKey].(map[string]any) + require.True(t, imageURLConverted) + assert.Equal(t, "data:image/png;base64,AQID", imageURL["url"]) + + anthropic, err := anthropicMessages([]llm.Message{message}) + require.NoError(t, err) + + anthropicContent, anthropicConverted := anthropic[0][jsonContentKey].([]map[string]any) + require.True(t, anthropicConverted) + + source, sourceConverted := anthropicContent[1]["source"].(map[string]any) + require.True(t, sourceConverted) + assert.Equal(t, "AQID", source["data"]) +} + +func TestMultipartProviderMappingsPreserveImageOnlyOrder(t *testing.T) { + t.Parallel() + + message := llm.Message{Metadata: nil, Role: llm.RoleUser, Content: []llm.Part{ + {Metadata: nil, ToolCall: nil, ToolResult: nil, Type: llm.PartImage, + Text: "", MIMEType: testImageMIME, Data: testImageData}, + {Metadata: nil, ToolCall: nil, ToolResult: nil, Type: llm.PartImage, + Text: "", MIMEType: "image/jpeg", Data: "BAUG"}, + }} + + responses, err := openAIResponseInput([]llm.Message{message}) + require.NoError(t, err) + + responseObject, responseObjectOK := responses[0].(map[string]any) + require.True(t, responseObjectOK) + + responseContent, responseContentOK := responseObject[jsonContentKey].([]map[string]any) + require.True(t, responseContentOK) + assert.Equal(t, []any{testImageDataURL, "data:image/jpeg;base64,BAUG"}, []any{ + responseContent[0][jsonImageURLKey], responseContent[1][jsonImageURLKey], + }) + + request := emptyCompletionRequest() + setTestRequestMessages(request, []llm.Message{message}) + chat, err := openAIChatMessages(request) + require.NoError(t, err) + + chatContent, chatContentOK := chat[0][jsonContentKey].([]map[string]any) + require.True(t, chatContentOK) + assert.Len(t, chatContent, 2) + + anthropic, err := anthropicMessages([]llm.Message{message}) + require.NoError(t, err) + + anthropicContent, anthropicContentOK := anthropic[0][jsonContentKey].([]map[string]any) + require.True(t, anthropicContentOK) + assert.Len(t, anthropicContent, 2) +} + +func TestMultipartProviderMappingsPreserveTextOnlyPayloads(t *testing.T) { + t.Parallel() + + message := llm.TextMessage(llm.RoleUser, "plain text") + responses, err := openAIResponseInput([]llm.Message{message}) + require.NoError(t, err) + + responseObject, responseObjectOK := responses[0].(map[string]any) + require.True(t, responseObjectOK) + + responseContent, responseContentOK := responseObject[jsonContentKey].([]map[string]any) + require.True(t, responseContentOK) + assert.Equal(t, "plain text", responseContent[0][jsonTextKey]) + + request := emptyCompletionRequest() + setTestRequestMessages(request, []llm.Message{message}) + chat, err := openAIChatMessages(request) + require.NoError(t, err) + assert.Equal(t, "plain text", chat[0][jsonContentKey]) + + anthropic, err := anthropicMessages([]llm.Message{message}) + require.NoError(t, err) + assert.Equal(t, "plain text", anthropic[0][jsonContentKey]) +} + +func TestMultipartProviderMappingsRejectMalformedAndUnsupportedRole(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + message llm.Message + }{ + {name: "empty data", message: testImageOnlyMessage(llm.RoleUser, testImageMIME, "")}, + {name: "malformed base64", message: testImageOnlyMessage(llm.RoleUser, testImageMIME, "not base64!")}, + {name: "unsupported MIME", message: testImageOnlyMessage(llm.RoleUser, "image/svg+xml", testImageData)}, + {name: "assistant role", message: testImageOnlyMessage(llm.RoleAssistant, testImageMIME, testImageData)}, + {name: "tool role", message: testImageOnlyMessage(llm.RoleTool, testImageMIME, testImageData)}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + for _, convert := range imageConverters(test.message) { + err := convert() + require.Error(t, err) + } + }) + } +} + +func TestCodexCompactionRejectsAssistantImagesBeforeFlattening(t *testing.T) { + t.Parallel() + + _, err := compactCodexResponseMessages([]llm.Message{ + testImageOnlyMessage(llm.RoleAssistant, testImageMIME, testImageData), + }) + require.Error(t, err) +} + +func imageConverters(message llm.Message) []func() error { + return []func() error{ + func() error { + _, err := openAIResponseInput([]llm.Message{message}) + + return err + }, + func() error { + request := emptyCompletionRequest() + setTestRequestMessages(request, []llm.Message{message}) + _, err := openAIChatMessages(request) + + return err + }, + func() error { + _, err := anthropicMessages([]llm.Message{message}) + + return err + }, + } +} diff --git a/internal/provider/messages.go b/internal/provider/messages.go index c34f8fe9..000841af 100644 --- a/internal/provider/messages.go +++ b/internal/provider/messages.go @@ -6,21 +6,35 @@ import ( "github.com/omarluq/librecode/internal/llm" ) -func openAIResponseInput(messages []llm.Message) []any { +func openAIResponseInput(messages []llm.Message) ([]any, error) { + if err := validateImageMessages(messages); err != nil { + return nil, err + } + input := []any{} for _, message := range messages { role, ok := openAIResponseInputRole(message.Role) - content := messageText(message) - if !ok || content == "" { + if !ok { + continue + } + + var content any = messageText(message) + if message.Role == llm.RoleUser { + if blocks := openAIResponseUserContent(message); len(blocks) > 0 { + content = blocks + } + } + + if content == "" { continue } input = append(input, map[string]any{jsonRoleKey: role, jsonContentKey: content}) } - return input + return input, nil } func openAIResponseInputRole(role llm.Role) (string, bool) { diff --git a/internal/provider/messages_internal_test.go b/internal/provider/messages_internal_test.go index a7f85b28..e5910ecd 100644 --- a/internal/provider/messages_internal_test.go +++ b/internal/provider/messages_internal_test.go @@ -4,6 +4,7 @@ import ( "testing" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/omarluq/librecode/internal/llm" ) @@ -14,8 +15,9 @@ func TestOpenAIResponseInputRoleMapping(t *testing.T) { request := emptyCompletionRequest() setTestRequestMessages(request, mixedReplayMessages()) - input := openAIResponseInput(request.Request.Messages) + input, err := openAIResponseInput(request.Request.Messages) + require.NoError(t, err) assert.Len(t, input, 7) for _, item := range input { @@ -26,6 +28,54 @@ func TestOpenAIResponseInputRoleMapping(t *testing.T) { } } +func TestOpenAIResponseInputMultipartPayloadShapes(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + message llm.Message + want []map[string]any + }{ + { + name: testImageWithText, + message: testImageMessage(), + want: []map[string]any{ + {jsonTypeKey: "input_text", jsonTextKey: testImagePrompt}, + {jsonTypeKey: "input_image", jsonImageURLKey: testImageDataURL}, + }, + }, + { + name: testImageOnly, + message: testImageOnlyMessage(llm.RoleUser, testImageMIME, testImageData), + want: []map[string]any{{ + jsonTypeKey: "input_image", jsonImageURLKey: testImageDataURL, + }}, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + input, err := openAIResponseInput([]llm.Message{ + test.message, + llm.TextMessage(llm.RoleAssistant, testAssistantReplay), + llm.TextMessage(llm.RoleTool, "tool output"), + }) + + require.NoError(t, err) + require.Len(t, input, 2) + user, ok := input[0].(map[string]any) + require.True(t, ok) + assert.JSONEq(t, jsonString(jsonUserRole), jsonString(user[jsonRoleKey])) + assert.Equal(t, test.want, user[jsonContentKey]) + assert.Equal(t, map[string]any{ + jsonRoleKey: jsonUserRole, jsonContentKey: testAssistantReplay, + }, input[1]) + }) + } +} + func TestOpenAIResponseInputRoleRejectsNonReplayRoles(t *testing.T) { t.Parallel() diff --git a/internal/provider/openai_chat.go b/internal/provider/openai_chat.go index b65a3bc8..837009b0 100644 --- a/internal/provider/openai_chat.go +++ b/internal/provider/openai_chat.go @@ -13,8 +13,13 @@ func (client *HTTPCompletionClient) completeOpenAIChat( ctx context.Context, request *CompletionRequest, ) (*llm.Response, error) { + messages, err := openAIChatMessages(request) + if err != nil { + return nil, err + } + state := openAIChatLoopState{ - messages: openAIChatMessages(request), + messages: messages, endpoint: joinEndpoint(request.Request.Model.BaseURL, "/chat/completions"), result: newResponse(), } @@ -182,7 +187,11 @@ func reasoningEffort(request *CompletionRequest) (string, bool) { return request.Request.ThinkingLevel, true } -func openAIChatMessages(request *CompletionRequest) []map[string]any { +func openAIChatMessages(request *CompletionRequest) ([]map[string]any, error) { + if err := validateImageMessages(request.Request.Messages); err != nil { + return nil, err + } + messages := []map[string]any{} if request.Request.SystemPrompt != "" { messages = append(messages, map[string]any{ @@ -194,15 +203,23 @@ func openAIChatMessages(request *CompletionRequest) []map[string]any { for _, message := range request.Request.Messages { role, ok := openAIRole(message.Role) - content := messageText(message) - if !ok || content == "" { + if !ok { + continue + } + + var content any = messageText(message) + if message.Role == llm.RoleUser { + content = openAIChatUserContent(message) + } + + if content == "" { continue } messages = append(messages, map[string]any{jsonRoleKey: role, jsonContentKey: content}) } - return messages + return messages, nil } func openAIChatAssistantToolMessage(result *providerResult) map[string]any { diff --git a/internal/provider/openai_chat_payload_internal_test.go b/internal/provider/openai_chat_payload_internal_test.go index af0e8c9f..f8695018 100644 --- a/internal/provider/openai_chat_payload_internal_test.go +++ b/internal/provider/openai_chat_payload_internal_test.go @@ -76,8 +76,9 @@ func TestOpenAIChatMessagesAndRoles(t *testing.T) { setTestRequestSystemPrompt(request, "system") setTestRequestMessages(request, mixedReplayMessages()) - messages := openAIChatMessages(request) + messages, err := openAIChatMessages(request) + require.NoError(t, err) assert.Len(t, messages, 8) assert.JSONEq(t, jsonString(jsonSystemRole), jsonString(messages[0][jsonRoleKey])) assert.JSONEq(t, jsonString(jsonUserRole), jsonString(messages[1][jsonRoleKey])) @@ -88,6 +89,57 @@ func TestOpenAIChatMessagesAndRoles(t *testing.T) { assert.Empty(t, mapped) } +func TestOpenAIChatMultipartPayloadShapes(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + message llm.Message + want []map[string]any + }{ + { + name: testImageWithText, + message: testImageMessage(), + want: []map[string]any{ + {jsonTypeKey: jsonTextKey, jsonTextKey: testImagePrompt}, + {jsonTypeKey: jsonImageURLKey, jsonImageURLKey: map[string]any{ + "url": testImageDataURL, + }}, + }, + }, + { + name: testImageOnly, + message: testImageOnlyMessage(llm.RoleUser, testImageMIME, testImageData), + want: []map[string]any{{ + jsonTypeKey: jsonImageURLKey, jsonImageURLKey: map[string]any{ + "url": testImageDataURL, + }, + }}, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + request := emptyCompletionRequest() + setTestRequestMessages(request, []llm.Message{ + test.message, + llm.TextMessage(llm.RoleAssistant, testAssistantReplay), + llm.TextMessage(llm.RoleTool, "tool output"), + }) + messages, err := openAIChatMessages(request) + + require.NoError(t, err) + require.Len(t, messages, 2) + assert.Equal(t, test.want, messages[0][jsonContentKey]) + assert.Equal(t, map[string]any{ + jsonRoleKey: jsonAssistantRole, jsonContentKey: testAssistantReplay, + }, messages[1]) + }) + } +} + func TestOpenAIChatPayloadAddsZAIStreamingOptions(t *testing.T) { t.Parallel() diff --git a/internal/provider/openai_responses.go b/internal/provider/openai_responses.go index a4bc38b0..010e95af 100644 --- a/internal/provider/openai_responses.go +++ b/internal/provider/openai_responses.go @@ -12,7 +12,11 @@ func (client *HTTPCompletionClient) completeOpenAIResponses( ctx context.Context, request *CompletionRequest, ) (*llm.Response, error) { - input := openAIResponseInput(request.Request.Messages) + input, err := openAIResponseInput(request.Request.Messages) + if err != nil { + return nil, err + } + endpoint := joinEndpoint(request.Request.Model.BaseURL, "/responses") return client.completeResponsesLoop(ctx, request, endpoint, openAIHeaders(request), input) @@ -22,12 +26,29 @@ func (client *HTTPCompletionClient) completeOpenAICodex( ctx context.Context, request *CompletionRequest, ) (*llm.Response, error) { - input := openAIResponseInput(compactResponseMessages(request.Request.Messages)) + messages, err := compactCodexResponseMessages(request.Request.Messages) + if err != nil { + return nil, err + } + + input, err := openAIResponseInput(messages) + if err != nil { + return nil, err + } + endpoint := joinEndpoint(request.Request.Model.BaseURL, "/codex/responses") return client.completeResponsesLoop(ctx, request, endpoint, codexHeaders(request), input) } +func compactCodexResponseMessages(messages []llm.Message) ([]llm.Message, error) { + if err := validateImageMessages(messages); err != nil { + return nil, err + } + + return compactResponseMessages(messages), nil +} + func (client *HTTPCompletionClient) completeResponsesLoop( ctx context.Context, request *CompletionRequest, diff --git a/internal/terminal/agent_tasks.go b/internal/terminal/agent_tasks.go index b7a3ee4c..af2dbe96 100644 --- a/internal/terminal/agent_tasks.go +++ b/internal/terminal/agent_tasks.go @@ -1098,7 +1098,7 @@ func (app *App) persistAgentCompletion(ctx context.Context, content string) { Role: database.RoleToolResult, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, }, &modelFacing, ) @@ -1563,13 +1563,17 @@ func (app *App) switchToAgentTaskSession( if app.restoreSessionView(sessionID) { promptHistory := app.promptHistory + promptHistoryImages := app.promptHistoryImages promptHistoryDraft := app.promptHistoryDraft + promptHistoryDraftImages := app.promptHistoryDraftImages promptHistoryIndex := app.promptHistoryIndex app.transcript.History = nil app.transcript.LineCache.reset() app.appendSessionMessages(messages) app.promptHistory = promptHistory + app.promptHistoryImages = promptHistoryImages app.promptHistoryDraft = promptHistoryDraft + app.promptHistoryDraftImages = promptHistoryDraftImages app.promptHistoryIndex = promptHistoryIndex messages = nil @@ -1703,15 +1707,18 @@ func (app *App) appendMissingSessionMessages(messages []database.SessionMessageE } app.appendMessage(chatMessage{ - CreatedAt: message.CreatedAt, - Role: role, - Content: message.Content, + CreatedAt: message.CreatedAt, + Role: role, + Content: message.Content, + Attachments: nil, }) appended = true if message.Role == database.RoleUser { - app.recordPromptHistory(message.Content) + app.recordPromptDraftHistory(promptDraft{ + Text: message.Content, Images: imageAttachmentsFromDatabase(message.Parts), + }) } } diff --git a/internal/terminal/agent_tasks_behavior_internal_test.go b/internal/terminal/agent_tasks_behavior_internal_test.go index f82955f3..ea130082 100644 --- a/internal/terminal/agent_tasks_behavior_internal_test.go +++ b/internal/terminal/agent_tasks_behavior_internal_test.go @@ -317,7 +317,7 @@ func TestAgentTaskCompletionFallbacksAndDelivery(t *testing.T) { assert.True(t, canceled) assert.Contains(t, app.statusMessage, "1 agent task(s) finished") require.Len(t, app.hiddenQueuedMessages, 1) - assert.Contains(t, app.hiddenQueuedMessages[0], "completion text") + assert.Contains(t, app.hiddenQueuedMessages[0].Text, "completion text") app.deliverAgentTaskCompletionText(t.Context(), "done", "duplicate") assert.Len(t, app.liveAgentCompletions, 1) @@ -391,7 +391,7 @@ func TestDiscoverRefreshAndTrackAgentTasks(t *testing.T) { app.deliverAgentTaskCompletion(t.Context(), &done) assert.Empty(t, app.agentTasks) require.Len(t, app.hiddenQueuedMessages, 1) - assert.Contains(t, app.hiddenQueuedMessages[0], "finished") + assert.Contains(t, app.hiddenQueuedMessages[0].Text, "finished") assert.True(t, stub.wasCanceled(behaviorRunning)) app.runtime = nil @@ -421,7 +421,7 @@ func TestRefreshActiveAgentTasksReconcilesMissedCompletion(t *testing.T) { assert.Empty(t, app.agentTasks) assert.True(t, stub.wasCanceled(behaviorRunning)) require.Len(t, app.hiddenQueuedMessages, 1) - assert.Contains(t, app.hiddenQueuedMessages[0], "finished after missed event") + assert.Contains(t, app.hiddenQueuedMessages[0].Text, "finished after missed event") } func TestRefreshActiveAgentTasksRetainsOmittedRunningTask(t *testing.T) { @@ -484,7 +484,7 @@ func TestTrackStartedTerminalTaskDeliversImmediately(t *testing.T) { app.working = true app.trackStartedAgentTask(t.Context(), agentToolEvent("", `{"task_id":"done"}`, false)) require.Len(t, app.liveAgentCompletions, 1) - assert.Contains(t, app.hiddenQueuedMessages[0], "immediate result") + assert.Contains(t, app.hiddenQueuedMessages[0].Text, "immediate result") } func TestAgentTaskStateAndRunningPredicates(t *testing.T) { @@ -651,7 +651,7 @@ func TestInspectAndLeaveAgentTaskSession(t *testing.T) { Role: database.RoleAssistant, Content: "loaded child transcript", Provider: "child-provider", - Model: "child-model", + Model: "child-model", Parts: nil, }) require.NoError(t, err) @@ -717,7 +717,7 @@ func TestLeaveAgentTaskSessionRefreshesDurableParentTranscript(t *testing.T) { Role: database.RoleAssistant, Content: "new durable parent message", Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) @@ -742,7 +742,7 @@ func TestRevisitAgentTaskSessionRefreshesDurableTranscript(t *testing.T) { Role: database.RoleAssistant, Content: "new durable child message", Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) @@ -1013,7 +1013,7 @@ func TestAgentTaskCompletionRoutesToParentWhileChildIsInspected(t *testing.T) { require.Len(t, app.liveAgentCompletions, 1) assert.Contains(t, app.liveAgentCompletions[0].Content, "parent completion") require.Len(t, app.hiddenQueuedMessages, 1) - assert.Contains(t, app.hiddenQueuedMessages[0], "parent completion") + assert.Contains(t, app.hiddenQueuedMessages[0].Text, "parent completion") } func TestActivePromptInspectionAllowsNestedTaskSelection(t *testing.T) { @@ -1419,7 +1419,7 @@ func TestInspectAgentTaskLoadFailureDoesNotSwitchSession(t *testing.T) { Role: database.RoleUser, Content: "child message", Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) }, diff --git a/internal/terminal/agent_tasks_live_internal_test.go b/internal/terminal/agent_tasks_live_internal_test.go index 11d7311c..2428bede 100644 --- a/internal/terminal/agent_tasks_live_internal_test.go +++ b/internal/terminal/agent_tasks_live_internal_test.go @@ -60,7 +60,7 @@ func TestInspectedAgentTaskRendersLiveStreamAndReloadsOnCompletion(t *testing.T) _, err := sessions.AppendMessage(t.Context(), child.ID, nil, &database.MessageEntity{ Timestamp: time.Now().UTC(), Role: database.RoleAssistant, Content: "finished result", - Provider: "provider", Model: "model", + Provider: "provider", Model: "model", Parts: nil, }) require.NoError(t, err) diff --git a/internal/terminal/app.go b/internal/terminal/app.go index 515c519d..adf8756c 100644 --- a/internal/terminal/app.go +++ b/internal/terminal/app.go @@ -50,9 +50,10 @@ const ( ) type chatMessage struct { - CreatedAt time.Time - Role transcript.Role - Content string + Attachments *attachmentSummaries + CreatedAt time.Time + Role transcript.Role + Content string } type activePromptState struct { @@ -61,6 +62,7 @@ type activePromptState struct { SessionID string UserEntryID string Prompt string + Images []imageAttachment ID uint64 Canceled bool } @@ -134,17 +136,18 @@ type RunOptions struct { // App is the terminal chat UI. type App struct { - lastEscape time.Time + extensionUI extui.State lastControlC time.Time workStartedAt time.Time + agentTasksRefreshedAt time.Time + lastEscape time.Time + workflows workflowInspector screen terminalScreen extensions extension.TerminalEventRunner - renderer *tui.Renderer - frame *tui.CellBuffer - lastResize *tcell.EventResize systemClipboard systemClipboardWriter - runtime *assistant.Runtime - workflows workflowInspector + imageClipboard systemClipboardImageReader + activeCompaction *activeCompactionState + activePrompt *activePromptState settings *database.DocumentRepository models *model.Registry auth *auth.Storage @@ -152,53 +155,57 @@ type App struct { keys *keybindings panel *panel.Model pendingParentID *string - activePrompt *activePromptState - activeCompaction *activeCompactionState + agentTaskWatches map[string]context.CancelFunc + workflowSteps map[string][]database.WorkflowAgentTaskDetail scopedEnabled map[string]bool - extensionUI extui.State + lastResize *tcell.EventResize + frame *tui.CellBuffer + workflowProgress map[string]workflowProgress + runtime *assistant.Runtime + sessionViews map[string]sessionViewState + deliveredAgentTasks map[string]struct{} + renderer *tui.Renderer theme terminalTheme - selectedPanelKind panel.Kind sessionID string - sessionViews map[string]sessionViewState - agentTaskSessionStack []string - agentTaskSummaryOwnerID string - statusMessage string - streamingText string streamingThinkingText string cwd string promptHistoryDraft string + promptHistoryDraftImages []imageAttachment mode appMode + streamingText string + statusMessage string + agentTaskSummaryOwnerID string + workflowPanelRunID string + workflowSummaryRunID string + selectedPanelKind panel.Kind resources core.ResourceSnapshot - transcript transcriptState runningToolBlocks []runningToolBlock - liveAgentCompletions []chatMessage + composerImages []imageAttachment agentTasks []database.AgentTaskEntity + liveAgentCompletions []chatMessage activeWorkflows []database.WorkflowRunEntity - workflowProgress map[string]workflowProgress - workflowSteps map[string][]database.WorkflowAgentTaskDetail - workflowSummaryRunID string - workflowPanelRunID string - agentTasksRefreshedAt time.Time - agentTaskWatches map[string]context.CancelFunc - deliveredAgentTasks map[string]struct{} - queuedMessages []string - hiddenQueuedMessages []string + agentTaskSessionStack []string + queuedMessages []promptDraft + hiddenQueuedMessages []promptDraft promptHistory []string + promptHistoryImages [][]imageAttachment scopedOrder []string composerBuffer tui.TextArea + transcript transcriptState tokenUsage model.TokenUsage selection mouseSelection transcriptList transcriptListSelection agentTaskSummarySelection agentTaskSummarySelection promptSequence uint64 + autocompleteSelection int workFrame int streamedToolEvents int escapePresses int promptHistoryIndex int scrollOffset int - autocompleteSelection int - autocompleteClosed bool sessionNamedOnly bool + autocompleteClosed bool + bracketedPaste bool hideThinking bool working bool compacting bool @@ -220,6 +227,8 @@ func Run(ctx context.Context, options *RunOptions) error { } screen.EnableMouse(tcell.MouseDragEvents) + + screen.EnablePaste() defer screen.Fini() app := newApp(screen, options) @@ -251,17 +260,26 @@ type terminalScreen interface { HideCursor() Show() SetClipboard(data []byte) + EnablePaste() ShowCursor(x, y int) Size() (width, height int) } func newApp(screen terminalScreen, options *RunOptions) *App { - app := &App{ + app := newAppState(screen, options) + app.addWelcomeMessage() + + return app +} + +func newAppState(screen terminalScreen, options *RunOptions) *App { + return &App{ screen: screen, renderer: tui.NewRenderer(screen), frame: nil, lastResize: nil, systemClipboard: newDesktopClipboard(), + imageClipboard: newDesktopClipboard(), runtime: options.Runtime, workflows: options.Workflows, extensions: options.Extensions, @@ -294,13 +312,17 @@ func newApp(screen terminalScreen, options *RunOptions) *App { agentTasksRefreshedAt: time.Time{}, agentTaskWatches: map[string]context.CancelFunc{}, deliveredAgentTasks: map[string]struct{}{}, - queuedMessages: []string{}, - hiddenQueuedMessages: []string{}, + queuedMessages: []promptDraft{}, + hiddenQueuedMessages: []promptDraft{}, promptHistory: []string{}, + promptHistoryImages: [][]imageAttachment{}, promptHistoryDraft: "", + promptHistoryDraftImages: nil, autocompleteSelection: 0, autocompleteClosed: false, composerBuffer: tui.NewTextArea(), + composerImages: []imageAttachment{}, + bracketedPaste: false, scopedOrder: []string{}, scopedEnabled: map[string]bool{}, sessionSortRecent: true, @@ -330,9 +352,6 @@ func newApp(screen terminalScreen, options *RunOptions) *App { streamingThinkingText: "", extensionUI: extui.NewState(), } - app.addWelcomeMessage() - - return app } func initialTranscriptState() transcriptState { @@ -703,13 +722,16 @@ func (app *App) appendSessionMessages(messages []database.SessionMessageEntity) for index := range messages { message := &messages[index] app.appendMessage(chatMessage{ - CreatedAt: message.CreatedAt, - Role: transcript.FromDatabaseRole(message.Role), - Content: message.Content, + CreatedAt: message.CreatedAt, + Role: transcript.FromDatabaseRole(message.Role), + Content: message.Content, + Attachments: databaseAttachmentSummaries(message.Parts), }) if message.Role == database.RoleUser { - app.recordPromptHistory(message.Content) + app.recordPromptDraftHistory(promptDraft{ + Text: message.Content, Images: imageAttachmentsFromDatabase(message.Parts), + }) } } } @@ -723,7 +745,7 @@ func (app *App) addMessage(role transcript.Role, content string) { } func newChatMessage(role transcript.Role, content string) chatMessage { - return chatMessage{CreatedAt: time.Now().UTC(), Role: role, Content: content} + return chatMessage{CreatedAt: time.Now().UTC(), Role: role, Content: content, Attachments: nil} } func emptyCachedRenderedMessage() cachedRenderedMessage { diff --git a/internal/terminal/async_events_internal_test.go b/internal/terminal/async_events_internal_test.go index fb4faa63..2aea3de3 100644 --- a/internal/terminal/async_events_internal_test.go +++ b/internal/terminal/async_events_internal_test.go @@ -525,9 +525,10 @@ func promptLifecycleEventCases() []promptLifecycleCase { app.streamingText = asyncTestPartial app.streamingThinkingText = "thought" app.transcript.Streaming.Blocks = []chatMessage{{ - Role: transcript.RoleAssistant, - Content: asyncTestPartial, - CreatedAt: time.Time{}, + Role: transcript.RoleAssistant, + Content: asyncTestPartial, + CreatedAt: time.Time{}, + Attachments: nil, }} app.runningToolBlocks = []runningToolBlock{{ StartedAt: time.Time{}, @@ -629,13 +630,13 @@ func TestApplyPromptErrorProcessesQueuedPrompt(t *testing.T) { app.activePrompt = newTestActivePrompt(nil) app.activePrompt.ID = 4 app.working = true - app.queuedMessages = []string{asyncTestQueuedText} + app.queuedMessages = promptDrafts(asyncTestQueuedText) app.applyPromptError(context.Background(), "provider failed", app.activePrompt.ID) waitForPromptRequest(t, client) assert.True(t, app.working) - assert.True(t, slices.Equal(app.queuedMessages, []string(nil))) + assert.True(t, slices.Equal(promptDraftTexts(app.queuedMessages), []string(nil))) } type streamEventApplyCase struct { diff --git a/internal/terminal/attachment_actions.go b/internal/terminal/attachment_actions.go new file mode 100644 index 00000000..d06552cb --- /dev/null +++ b/internal/terminal/attachment_actions.go @@ -0,0 +1,60 @@ +package terminal + +import "github.com/gdamore/tcell/v3" + +func (app *App) handleAttachmentKey(event *tcell.EventKey) bool { + switch { + case app.keys.matches(event, actionAttachmentPasteImage): + return app.pasteClipboardImage() + case app.keys.matches(event, actionAttachmentRemoveLast): + if len(app.composerImages) == 0 { + return false + } + + app.composerImages = app.composerImages[:len(app.composerImages)-1] + app.setStatus("removed last image attachment") + + return true + case app.keys.matches(event, actionAttachmentClear): + if len(app.composerImages) == 0 { + return false + } + + app.composerImages = nil + app.setStatus("cleared image attachments") + + return true + default: + return false + } +} + +func (app *App) pasteClipboardImage() bool { + if app.imageClipboard == nil { + return false + } + + data, err := app.imageClipboard.ReadImage() + if err != nil { + app.setStatus(err.Error()) + + return true + } + // An empty image clipboard is deliberately not consumed: the terminal may + // immediately deliver an ordinary bracketed text paste for the same key. + if len(data) == 0 { + return false + } + + attachment, err := validateClipboardPNG(data, app.composerImages) + if err != nil { + app.setStatus(err.Error()) + + return true + } + + app.composerImages = append(app.composerImages, attachment) + app.setStatus("attached " + attachment.Name) + + return true +} diff --git a/internal/terminal/attachments.go b/internal/terminal/attachments.go new file mode 100644 index 00000000..158abab2 --- /dev/null +++ b/internal/terminal/attachments.go @@ -0,0 +1,249 @@ +package terminal + +import ( + "bytes" + "errors" + "fmt" + "image" + _ "image/png" // Register PNG for clipboard image configuration validation. + "slices" + "strings" + + "github.com/omarluq/librecode/internal/assistant" + "github.com/omarluq/librecode/internal/database" + "github.com/omarluq/librecode/internal/model" +) + +const ( + clipboardImageMIME = "image/png" + maxComposerImages = 4 + maxComposerImageBytes = 5 << 20 + maxComposerImageTotal = 20 << 20 + maxComposerImagePixels = 40_000_000 +) + +type imageAttachment struct { + Name string + MIMEType string + Data []byte + Width int + Height int +} + +type attachmentSummaries []attachmentSummary + +type attachmentSummary struct { + Name string + MIMEType string + Width int + Height int + Size int +} + +type promptDraft struct { + Text string + Images []imageAttachment +} + +func (draft promptDraft) empty() bool { + return strings.TrimSpace(draft.Text) == "" && len(draft.Images) == 0 +} + +func cloneImageAttachments(images []imageAttachment) []imageAttachment { + cloned := make([]imageAttachment, len(images)) + for index := range images { + cloned[index] = images[index] + cloned[index].Data = bytes.Clone(images[index].Data) + } + + return cloned +} + +func clonePromptDraft(draft promptDraft) promptDraft { + return promptDraft{Text: draft.Text, Images: cloneImageAttachments(draft.Images)} +} + +func cloneImageAttachmentGroups(groups [][]imageAttachment) [][]imageAttachment { + cloned := make([][]imageAttachment, len(groups)) + for index := range groups { + cloned[index] = cloneImageAttachments(groups[index]) + } + + return cloned +} + +func clonePromptDrafts(drafts []promptDraft) []promptDraft { + cloned := make([]promptDraft, len(drafts)) + for index := range drafts { + cloned[index] = clonePromptDraft(drafts[index]) + } + + return cloned +} + +func (draft promptDraft) assistantImages() []assistant.ImageAttachment { + images := make([]assistant.ImageAttachment, len(draft.Images)) + for index, item := range draft.Images { + images[index] = assistant.ImageAttachment{ + Name: item.Name, MIMEType: item.MIMEType, Data: bytes.Clone(item.Data), + Width: item.Width, Height: item.Height, + } + } + + return images +} + +func summarizeAttachments(images []imageAttachment) *attachmentSummaries { + result := make([]attachmentSummary, len(images)) + for index, item := range images { + result[index] = attachmentSummary{ + Name: item.Name, MIMEType: item.MIMEType, Width: item.Width, + Height: item.Height, Size: len(item.Data), + } + } + + summaries := attachmentSummaries(result) + + return &summaries +} + +func imageAttachmentsFromDatabase(parts []database.MessagePartEntity) []imageAttachment { + images := make([]imageAttachment, 0, len(parts)) + for index := range parts { + part := &parts[index] + if part.Type == database.MessagePartImage { + images = append(images, imageAttachment{ + Name: part.Name, MIMEType: part.MIMEType, Data: bytes.Clone(part.Data), + Width: part.Width, Height: part.Height, + }) + } + } + + return images +} + +func databaseAttachmentSummaries(parts []database.MessagePartEntity) *attachmentSummaries { + result := make([]attachmentSummary, 0, len(parts)) + for _, part := range parts { + if part.Type == database.MessagePartImage { + result = append(result, attachmentSummary{ + Name: part.Name, MIMEType: part.MIMEType, Width: part.Width, + Height: part.Height, Size: len(part.Data), + }) + } + } + + summaries := attachmentSummaries(result) + + return &summaries +} + +func (app *App) composerDraftEmpty() bool { + return app.composerBuffer.Empty() && len(app.composerImages) == 0 +} + +func (app *App) currentDraft() promptDraft { + return promptDraft{ + Text: strings.TrimSpace(app.composerBuffer.TextValue()), + Images: cloneImageAttachments(app.composerImages), + } +} + +func (app *App) consumeDraft() promptDraft { + draft := promptDraft{Text: strings.TrimSpace(app.composerBuffer.Clear()), Images: app.composerImages} + app.composerImages = nil + + return draft +} + +func (app *App) restoreDraft(draft promptDraft) { + app.composerBuffer.SetText(draft.Text) + app.composerImages = cloneImageAttachments(draft.Images) +} + +func validateClipboardPNG(data []byte, existing []imageAttachment) (imageAttachment, error) { + if err := validateClipboardImageSize(data, existing); err != nil { + return imageAttachment{}, err + } + + return decodeClipboardPNG(data, len(existing)+1) +} + +func validateClipboardImageSize(data []byte, existing []imageAttachment) error { + if len(existing) >= maxComposerImages { + return fmt.Errorf("image attachment limit is %d", maxComposerImages) + } + + if len(data) > maxComposerImageBytes { + return errors.New("image exceeds the 5 MiB limit") + } + + if len(data) == 0 { + return errors.New("clipboard contains no image") + } + + total := len(data) + for _, item := range existing { + total += len(item.Data) + } + + if total > maxComposerImageTotal { + return errors.New("image attachments exceed the 20 MiB limit") + } + + return nil +} + +func decodeClipboardPNG(data []byte, sequence int) (imageAttachment, error) { + config, format, err := image.DecodeConfig(bytes.NewReader(data)) + if err != nil { + return imageAttachment{}, fmt.Errorf("decode clipboard image: %w", err) + } + + if format != "png" { + return imageAttachment{}, errors.New("clipboard image must be PNG") + } + + if config.Width <= 0 || config.Height <= 0 || config.Width > maxComposerImagePixels/config.Height { + return imageAttachment{}, errors.New("image dimensions must be positive and at most 40 megapixels") + } + + return imageAttachment{ + Name: fmt.Sprintf("paste-%d.png", sequence), MIMEType: clipboardImageMIME, Data: bytes.Clone(data), + Width: config.Width, Height: config.Height, + }, nil +} + +func (app *App) selectedModelSupportsImages() bool { + if app.models == nil { + return true + } + + models := app.models.All() + for index := range models { + candidate := &models[index] + if candidate.Provider == app.currentProvider() && candidate.ID == app.currentModel() { + return slices.Contains(candidate.Input, model.InputImage) + } + } + + return true +} + +func (app *App) validateDraftModel(draft promptDraft) bool { + if len(draft.Images) == 0 || app.selectedModelSupportsImages() { + return true + } + + app.setStatus("selected model does not support image input") + + return false +} + +func derefAttachmentSummaries(summaries *attachmentSummaries) []attachmentSummary { + if summaries == nil { + return nil + } + + return *summaries +} diff --git a/internal/terminal/attachments_internal_test.go b/internal/terminal/attachments_internal_test.go new file mode 100644 index 00000000..895905d0 --- /dev/null +++ b/internal/terminal/attachments_internal_test.go @@ -0,0 +1,241 @@ +package terminal + +import ( + "bytes" + "context" + "image" + "image/png" + "strings" + "testing" + + "github.com/gdamore/tcell/v3" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/omarluq/librecode/internal/model" +) + +const testImageAttachmentName = "image.png" + +type imageClipboardStub struct { + err error + data []byte +} + +func (stub imageClipboardStub) ReadImage() ([]byte, error) { return stub.data, stub.err } + +func testPNG(t *testing.T, width, height int) []byte { + t.Helper() + + var buffer bytes.Buffer + require.NoError(t, png.Encode(&buffer, image.NewRGBA(image.Rect(0, 0, width, height)))) + + return buffer.Bytes() +} + +func TestClipboardImagePasteValidateRemoveAndClear(t *testing.T) { + t.Parallel() + app := newRenderTestApp(t) + data := testPNG(t, 3, 2) + app.imageClipboard = imageClipboardStub{err: nil, data: data} + + assert.True(t, app.pasteClipboardImage()) + require.Len(t, app.composerImages, 1) + assert.Equal(t, "paste-1.png", app.composerImages[0].Name) + assert.Equal(t, 3, app.composerImages[0].Width) + assert.Equal(t, "attached paste-1.png", app.statusMessage) + + data[0] ^= 0xff + assert.NotEqual(t, data[0], app.composerImages[0].Data[0]) + + assert.True(t, app.handleAttachmentKey(tcell.NewEventKey(tcell.KeyRune, "r", tcell.ModAlt))) + assert.Empty(t, app.composerImages) + app.composerImages = []imageAttachment{ + {Name: "first-image", MIMEType: "", Data: nil, Width: 0, Height: 0}, + {Name: "second-image", MIMEType: "", Data: nil, Width: 0, Height: 0}, + } + assert.True(t, app.handleAttachmentKey(tcell.NewEventKey(tcell.KeyRune, "c", tcell.ModAlt))) + assert.Empty(t, app.composerImages) +} + +func TestClipboardImageFailuresPreserveDraft(t *testing.T) { + t.Parallel() + app := newRenderTestApp(t) + app.composerImages = []imageAttachment{{Name: "existing", MIMEType: "", Data: []byte{1}, Width: 0, Height: 0}} + app.imageClipboard = imageClipboardStub{err: nil, data: []byte("not png")} + assert.True(t, app.pasteClipboardImage()) + assert.Equal(t, "existing", app.composerImages[0].Name) + assert.Contains(t, app.statusMessage, "decode clipboard image") + + app.imageClipboard = imageClipboardStub{err: nil, data: nil} + assert.False(t, app.pasteClipboardImage()) + require.Len(t, app.composerImages, 1) +} + +func TestClipboardImageLimits(t *testing.T) { + t.Parallel() + pngData := testPNG(t, 1, 1) + _, err := validateClipboardPNG(make([]byte, maxComposerImageBytes+1), nil) + require.ErrorContains(t, err, "5 MiB") + _, err = validateClipboardPNG(pngData, make([]imageAttachment, maxComposerImages)) + require.ErrorContains(t, err, "limit is 4") + + existing := imageAttachment{ + Name: "", MIMEType: "", Data: make([]byte, maxComposerImageTotal), Width: 0, Height: 0, + } + _, err = validateClipboardPNG(pngData, []imageAttachment{existing}) + require.ErrorContains(t, err, "20 MiB") +} + +func TestBracketedPasteInsertsLiterally(t *testing.T) { + t.Parallel() + app := newRenderTestApp(t) + ctx := context.Background() + _, err := app.handleEvent(ctx, tcell.NewEventPaste(true)) + require.NoError(t, err) + + for _, event := range []*tcell.EventKey{ + tcell.NewEventKey(tcell.KeyRune, "/", tcell.ModNone), + tcell.NewEventKey(tcell.KeyEnter, "", tcell.ModNone), + tcell.NewEventKey(tcell.KeyTab, "", tcell.ModNone), + tcell.NewEventKey(tcell.KeyRune, "x", tcell.ModNone), + } { + _, err = app.handleEvent(ctx, event) + require.NoError(t, err) + } + + _, err = app.handleEvent(ctx, tcell.NewEventPaste(false)) + require.NoError(t, err) + assert.Equal(t, "/\n\tx", app.composerBuffer.TextValue()) + assert.False(t, app.working) +} + +func TestPromptDraftAndSessionViewCloneImageBytes(t *testing.T) { + t.Parallel() + app := newRenderTestApp(t) + app.sessionID = "session" + app.composerImages = []imageAttachment{{Name: "image", MIMEType: "", Data: []byte{1, 2}, Width: 0, Height: 0}} + queuedImage := imageAttachment{Name: "", MIMEType: "", Data: []byte{3}, Width: 0, Height: 0} + app.queuedMessages = []promptDraft{{Text: "queued", Images: []imageAttachment{queuedImage}}} + app.saveSessionView() + app.composerImages[0].Data[0] = 9 + app.queuedMessages[0].Images[0].Data[0] = 9 + view := app.sessionViews["session"] + assert.Equal(t, byte(1), view.composerImages[0].Data[0]) + assert.Equal(t, byte(3), view.queuedMessages[0].Images[0].Data[0]) +} + +func TestDraftModelCapabilityValidation(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + app.models = nil + draft := promptDraft{Text: "", Images: []imageAttachment{{ + Name: "", MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1, + }}} + assert.True(t, app.validateDraftModel(draft)) + + app.models = model.NewRegistry(&model.RegistryOptions{ + ConfigReader: nil, Auth: nil, ModelsPath: "", + BuiltIns: []model.Model{terminalCapabilityTestModel(app.currentProvider(), app.currentModel())}, + Discovery: disabledModelDiscovery(), + }) + assert.False(t, app.validateDraftModel(draft)) + assert.Equal(t, "selected model does not support image input", app.statusMessage) +} + +func terminalCapabilityTestModel(provider, id string) model.Model { + return model.Model{ + ThinkingLevelMap: nil, Headers: nil, Compat: nil, Provider: provider, ID: id, + Name: id, API: "", BaseURL: "", Input: []model.InputMode{model.InputText}, + Cost: model.Cost{Input: 0, Output: 0, CacheRead: 0, CacheWrite: 0}, + ContextWindow: 0, MaxTokens: 0, Reasoning: false, + } +} + +func TestImageOnlyDraftIsNotTreatedAsEmptyAndComposerLayoutIsBounded(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + app.composerImages = []imageAttachment{ + {Name: "image-a", MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1}, + {Name: "image-b", MIMEType: clipboardImageMIME, Data: []byte{2}, Width: 1, Height: 1}, + {Name: "image-c", MIMEType: clipboardImageMIME, Data: []byte{3}, Width: 1, Height: 1}, + } + assert.False(t, app.composerDraftEmpty()) + rendered := app.renderComposerEditor(40, 2) + assert.LessOrEqual(t, len(rendered.Lines), 2+composerBorderRows) + assert.Contains(t, rendered.Lines[1].Text, "3 more attachments") + + app.handleEscapePresses(t.Context(), 1) + assert.True(t, app.composerDraftEmpty()) + assert.Equal(t, "editor cleared", app.statusMessage) +} + +func TestControlCClearsImageOnlyDraftBeforeExit(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + app.composerImages = []imageAttachment{{ + Name: testImageAttachmentName, MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1, + }} + + shouldQuit, err := app.handleKey( + t.Context(), tcell.NewEventKey(tcell.KeyCtrlC, "", tcell.ModCtrl), + ) + require.NoError(t, err) + assert.False(t, shouldQuit) + assert.True(t, app.composerDraftEmpty()) +} + +func TestReadOnlyInspectionBlocksImageAndBracketedPaste(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + app.sessionID = "child" + app.activePrompt = newTestActivePrompt(func() {}) + app.activePrompt.SessionID = sessionCommandsParentID + app.working = true + app.imageClipboard = imageClipboardStub{err: nil, data: testPNG(t, 1, 1)} + + app.bracketedPaste = true + _, err := app.handleEvent( + t.Context(), tcell.NewEventKey(tcell.KeyRune, "x", tcell.ModNone), + ) + require.NoError(t, err) + assert.False(t, app.bracketedPaste) + assert.Empty(t, app.composerBuffer.TextValue()) + + _, err = app.handleEvent(t.Context(), tcell.NewEventPaste(true)) + require.NoError(t, err) + assert.False(t, app.bracketedPaste) + + shouldQuit, err := app.handleKey( + t.Context(), tcell.NewEventKey(tcell.KeyRune, "v", tcell.ModCtrl), + ) + require.NoError(t, err) + assert.False(t, shouldQuit) + assert.Empty(t, app.composerImages) + assert.Equal(t, readOnlyAgentInspectionStatus, app.statusMessage) +} + +func TestAttachmentRenderingContainsMetadataNotData(t *testing.T) { + t.Parallel() + app := newRenderTestApp(t) + secret := []byte("RAW-SECRET") + attachment := imageAttachment{Name: "paste-1.png", MIMEType: "image/png", Data: secret, Width: 10, Height: 20} + app.composerImages = []imageAttachment{attachment} + composer := app.attachmentChipLines(80) + assert.Contains(t, composer[0].Text, "paste-1.png") + assert.NotContains(t, composer[0].Text, string(secret)) + + lines := app.renderUserMessage(80, "", *summarizeAttachments([]imageAttachment{attachment})) + + var rendered strings.Builder + for _, line := range lines { + rendered.WriteString(line.Text) + } + + assert.Contains(t, rendered.String(), "10×20") + assert.NotContains(t, rendered.String(), string(secret)) +} diff --git a/internal/terminal/auth_commands_internal_test.go b/internal/terminal/auth_commands_internal_test.go index 8ac49cad..a9b25c30 100644 --- a/internal/terminal/auth_commands_internal_test.go +++ b/internal/terminal/auth_commands_internal_test.go @@ -176,7 +176,7 @@ func TestParentIDFromEntry(t *testing.T) { Role: "", Content: "", Provider: "", - Model: "", + Model: "", Parts: nil, }, Summary: "", ToolStatus: "", diff --git a/internal/terminal/clipboard.go b/internal/terminal/clipboard.go index f62e7aca..62286dfd 100644 --- a/internal/terminal/clipboard.go +++ b/internal/terminal/clipboard.go @@ -16,10 +16,15 @@ type systemClipboardWriter interface { WriteText(text string) error } +type systemClipboardImageReader interface { + ReadImage() ([]byte, error) +} + type desktopClipboard struct { prepare func() error init func() error write func(clipboard.Format, []byte) <-chan struct{} + read func(clipboard.Format) []byte } func newDesktopClipboard() desktopClipboard { @@ -27,6 +32,7 @@ func newDesktopClipboard() desktopClipboard { prepare: defaultPrepareDesktopClipboardEnvironment, init: clipboard.Init, write: clipboard.Write, + read: clipboard.Read, } } @@ -82,6 +88,20 @@ func candidateWaylandDisplay( return "" } +func (writer desktopClipboard) ReadImage() ([]byte, error) { + if writer.prepare != nil { + if err := writer.prepare(); err != nil { + return nil, terminalError(err, "prepare system clipboard") + } + } + + if err := writer.init(); err != nil { + return nil, terminalError(err, "init system clipboard") + } + + return append([]byte(nil), writer.read(clipboard.FmtImage)...), nil +} + func (writer desktopClipboard) WriteText(text string) error { if text == "" { return nil diff --git a/internal/terminal/clipboard_internal_test.go b/internal/terminal/clipboard_internal_test.go index 3e840681..c6f2d04c 100644 --- a/internal/terminal/clipboard_internal_test.go +++ b/internal/terminal/clipboard_internal_test.go @@ -429,6 +429,7 @@ func callDesktopClipboard(text string, initErr error, changed <-chan struct{}) d return initErr }, + read: func(clipboard.Format) []byte { return nil }, write: func(_ clipboard.Format, data []byte) <-chan struct{} { result.writes = append(result.writes, append([]byte(nil), data...)) diff --git a/internal/terminal/compact_commands_internal_test.go b/internal/terminal/compact_commands_internal_test.go index d18d189f..64e00c5d 100644 --- a/internal/terminal/compact_commands_internal_test.go +++ b/internal/terminal/compact_commands_internal_test.go @@ -181,7 +181,7 @@ func TestHandleCompactDoneStartsQueuedPrompt(t *testing.T) { app := newPromptSendTestApp(t, client) app.compacting = true app.activeCompaction = &activeCompactionState{Cancel: func() {}, ID: 9, QueuedStart: 0} - app.queuedMessages = []string{"queued after compact"} + app.queuedMessages = promptDrafts("queued after compact") app.applyCompactDone(context.Background(), &asyncEvent{ Response: nil, ToolCallEvent: nil, ToolEvent: nil, Usage: &model.TokenUsage{ @@ -393,7 +393,7 @@ func appendTerminalCompactMessage( Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, } entry, err := app.runtime.SessionRepository().AppendMessage(context.Background(), sessionID, parentID, message) @@ -524,7 +524,7 @@ func TestCompactErrorRestoresQueuedPrompt(t *testing.T) { t.Parallel() app := newRenderTestApp(t) - app.queuedMessages = []string{"preexisting", "during compaction"} + app.queuedMessages = promptDrafts("preexisting", "during compaction") app.compacting = true app.activeCompaction = &activeCompactionState{Cancel: func() {}, ID: 9, QueuedStart: 1} @@ -534,7 +534,7 @@ func TestCompactErrorRestoresQueuedPrompt(t *testing.T) { }) assert.Equal(t, "during compaction", app.composerBuffer.TextValue()) - assert.Equal(t, []string{"preexisting"}, app.queuedMessages) + assert.Equal(t, []string{"preexisting"}, promptDraftTexts(app.queuedMessages)) } func TestCompactFormattingHelpers(t *testing.T) { diff --git a/internal/terminal/extension_events_internal_test.go b/internal/terminal/extension_events_internal_test.go index 33c0815f..f5d40af0 100644 --- a/internal/terminal/extension_events_internal_test.go +++ b/internal/terminal/extension_events_internal_test.go @@ -79,6 +79,50 @@ end) } } +func TestExtensionPromptSubmitMutationPreservesImages(t *testing.T) { + t.Parallel() + + app := newExtensionRuntimeTestApp(t, ` +librecode.on("prompt_submit", function() + librecode.buf.set_text("composer", "mutated by extension") +end) +`) + app.working = true + app.composerBuffer.SetText("original") + app.composerImages = []imageAttachment{{ + Name: testImageAttachmentName, MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1, + }} + + shouldQuit, err := app.submit(t.Context()) + require.NoError(t, err) + assert.False(t, shouldQuit) + require.Len(t, app.queuedMessages, 1) + assert.Equal(t, "mutated by extension", app.queuedMessages[0].Text) + require.Len(t, app.queuedMessages[0].Images, 1) + assert.Equal(t, []byte{1}, app.queuedMessages[0].Images[0].Data) +} + +func TestExtensionPromptSubmitRevalidatesImageDraft(t *testing.T) { + t.Parallel() + + app := newExtensionRuntimeTestApp(t, ` +librecode.on("prompt_submit", function() + librecode.buf.set_text("composer", "/model") +end) +`) + app.composerBuffer.SetText("describe this") + app.composerImages = []imageAttachment{{ + Name: testImageAttachmentName, MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1, + }} + + shouldQuit, err := app.submit(t.Context()) + require.NoError(t, err) + assert.False(t, shouldQuit) + assert.Equal(t, "/model", app.composerBuffer.TextValue()) + assert.Len(t, app.composerImages, 1) + assert.Equal(t, "slash commands do not accept image attachments", app.statusMessage) +} + func TestExtensionRuntimeBuffersPersistBetweenEvents(t *testing.T) { t.Parallel() diff --git a/internal/terminal/input.go b/internal/terminal/input.go index 8d7a8b2e..100a5ac1 100644 --- a/internal/terminal/input.go +++ b/internal/terminal/input.go @@ -16,7 +16,31 @@ func (app *App) handleEvent(ctx context.Context, event tcell.Event) (bool, error switch typedEvent := event.(type) { case *tcell.EventResize: return false, app.applyResizeEvent(ctx, typedEvent) + case *tcell.EventPaste: + if app.inspectingWhilePromptRuns() { + app.bracketedPaste = false + app.setStatus(readOnlyAgentInspectionStatus) + + return false, nil + } + + app.bracketedPaste = typedEvent.Start() + + return false, nil case *tcell.EventKey: + if app.bracketedPaste { + if app.inspectingWhilePromptRuns() { + app.bracketedPaste = false + app.setStatus(readOnlyAgentInspectionStatus) + + return false, nil + } + + app.insertPastedKey(typedEvent) + + return false, nil + } + return app.handleKey(ctx, typedEvent) case *tcell.EventMouse: app.handleMouse(typedEvent) @@ -55,6 +79,10 @@ func (app *App) handlePriorityKey(ctx context.Context, event *tcell.EventKey) ke return result } + if app.handleAttachmentKey(event) { + return keyHandlingResult{err: nil, shouldQuit: false, handled: true} + } + if app.handleAutocompletePriorityKey(event) { return keyHandlingResult{err: nil, shouldQuit: false, handled: true} } @@ -166,7 +194,7 @@ func (app *App) handleInlineListsAndExtensionKey( } func (app *App) handleForceExitKey(event *tcell.EventKey) (handled, shouldQuit bool) { - if !app.keys.matches(event, actionForceExit) || !app.composerBuffer.Empty() { + if !app.keys.matches(event, actionForceExit) || !app.composerDraftEmpty() { return false, false } @@ -258,9 +286,24 @@ func (app *App) handleFocusedAutocompleteKey(event *tcell.EventKey) bool { return app.handleAutocompleteKey(event) } +func (app *App) insertPastedKey(event *tcell.EventKey) { + if event.Key() == tcell.KeyRune { + app.composerBuffer.InsertRune(tui.EventRune(event)) + } + + if event.Key() == tcell.KeyEnter { + app.composerBuffer.InsertRune('\n') + } + + if event.Key() == tcell.KeyTab { + app.composerBuffer.InsertRune('\t') + } +} + func (app *App) handleInputKey(ctx context.Context, event *tcell.EventKey) (bool, error) { - if app.keys.matches(event, actionInputClear) && !app.composerBuffer.Empty() { + if app.keys.matches(event, actionInputClear) && !app.composerDraftEmpty() { app.composerBuffer.Clear() + app.composerImages = nil app.resetPromptHistoryNavigation() app.resetAutocompleteSelection() app.escapePresses = 0 diff --git a/internal/terminal/input_escape.go b/internal/terminal/input_escape.go index c6e0606c..6d2c37a6 100644 --- a/internal/terminal/input_escape.go +++ b/internal/terminal/input_escape.go @@ -30,8 +30,9 @@ func (app *App) handleEscapePresses(ctx context.Context, presses int) { } app.escapePresses = 0 - if !app.composerBuffer.Empty() { + if !app.composerDraftEmpty() { app.composerBuffer.Clear() + app.composerImages = nil app.resetPromptHistoryNavigation() app.setStatus("editor cleared") diff --git a/internal/terminal/interrupt_internal_test.go b/internal/terminal/interrupt_internal_test.go index dc8478f8..b382c67a 100644 --- a/internal/terminal/interrupt_internal_test.go +++ b/internal/terminal/interrupt_internal_test.go @@ -187,6 +187,7 @@ func newTestActivePrompt(cancel context.CancelFunc) *activePromptState { ParentEntryID: nil, SessionID: "", UserEntryID: "", + Images: nil, Prompt: interruptTestPrompt, ID: 1, Canceled: false, diff --git a/internal/terminal/keybindings.go b/internal/terminal/keybindings.go index 6a1aab6d..c707cb3d 100644 --- a/internal/terminal/keybindings.go +++ b/internal/terminal/keybindings.go @@ -31,6 +31,9 @@ const ( actionInputNewLine actionID = "tui.tui.newLine" actionInputSubmit actionID = "tui.tui.submit" actionInputTab actionID = "tui.tui.tab" + actionAttachmentPasteImage actionID = "tui.attachment.pasteImage" + actionAttachmentRemoveLast actionID = "tui.attachment.removeLast" + actionAttachmentClear actionID = "tui.attachment.clear" actionSelectUp actionID = "tui.select.up" actionSelectDown actionID = "tui.select.down" actionSelectPageUp actionID = "tui.select.pageUp" @@ -109,6 +112,9 @@ func defaultKeybindingDefinitions() map[actionID]keyBindingDefinition { actionInputNewLine: binding("Insert newline", "shift+enter"), actionInputSubmit: binding("Submit input", "enter"), actionInputTab: binding("Autocomplete or toggle scope", "tab"), + actionAttachmentPasteImage: binding("Paste clipboard image", "ctrl+v", "super+v"), + actionAttachmentRemoveLast: binding("Remove last image attachment", "alt+r"), + actionAttachmentClear: binding("Clear image attachments", "alt+c"), actionSelectUp: binding("Move selection up", "up"), actionSelectDown: binding("Move selection down", "down"), actionSelectPageUp: binding("Page selection up", "pageUp"), diff --git a/internal/terminal/message_render.go b/internal/terminal/message_render.go index 900c5ab4..7cf9a3e9 100644 --- a/internal/terminal/message_render.go +++ b/internal/terminal/message_render.go @@ -1,6 +1,7 @@ package terminal import ( + "fmt" "strings" "github.com/gdamore/tcell/v3" @@ -68,7 +69,7 @@ func (app *App) renderMessageDetailed(width int, message chatMessage) cachedRend case transcript.RoleAssistant: panic("assistant message handled before role switch") case transcript.RoleUser: - lines = app.renderUserMessage(width, message.Content) + lines = app.renderUserMessage(width, message.Content, derefAttachmentSummaries(message.Attachments)) case transcript.RoleToolResult, transcript.RoleBashExecution: lines = app.renderToolMessage(width, message) case transcript.RoleThinking: @@ -84,9 +85,14 @@ func (app *App) renderMessageDetailed(width int, message chatMessage) cachedRend return cachedRenderedMessage{Lines: lines, ListItems: []markdownListItemRange{}, Valid: true} } -func (app *App) renderUserMessage(width int, content string) []tui.Line { +func (app *App) renderUserMessage(width int, content string, attachments []attachmentSummary) []tui.Line { innerWidth := max(1, width-messageBoxHorizontalPadding) + wrapped := tui.Wrap(content, innerWidth) + for _, item := range attachments { + wrapped = append(wrapped, tui.Truncate(attachmentSummaryText(item), innerWidth)) + } + lines := make([]tui.Line, 0, len(wrapped)+defaultMessageExtraRows) lines = append(lines, @@ -115,7 +121,12 @@ func (app *App) renderQueuedMessages(width int) []tui.Line { header := "queued follow-up " + tui.Int(index+1) lines = append(lines, tui.NewLine(style.Bold(true), " "+tui.PadRight(header, innerWidth)+" ")) - for _, line := range tui.Wrap(message, innerWidth) { + for _, line := range tui.Wrap(message.Text, innerWidth) { + lines = append(lines, tui.NewLine(style, " "+tui.PadRight(line, innerWidth)+" ")) + } + + for _, item := range *summarizeAttachments(message.Images) { + line := tui.Truncate(attachmentSummaryText(item), innerWidth) lines = append(lines, tui.NewLine(style, " "+tui.PadRight(line, innerWidth)+" ")) } } @@ -125,6 +136,15 @@ func (app *App) renderQueuedMessages(width int) []tui.Line { return lines } +func attachmentSummaryText(item attachmentSummary) string { + format := strings.ToUpper(strings.TrimPrefix(item.MIMEType, "image/")) + + return fmt.Sprintf( + "[image: %s · %s · %d×%d · %s]", + item.Name, format, item.Width, item.Height, formatByteSize(item.Size), + ) +} + func (app *App) renderStreamingMessage(width int, content string) []tui.Line { wrapped := tui.Wrap(strings.TrimSpace(content), width) style := app.theme.style(colorText) diff --git a/internal/terminal/model_test_helpers_internal_test.go b/internal/terminal/model_test_helpers_internal_test.go index cbc2896b..bde688e7 100644 --- a/internal/terminal/model_test_helpers_internal_test.go +++ b/internal/terminal/model_test_helpers_internal_test.go @@ -12,3 +12,21 @@ func disabledModelDiscovery() model.DiscoveryOptions { Enabled: false, } } + +func promptDrafts(texts ...string) []promptDraft { + drafts := make([]promptDraft, len(texts)) + for index, text := range texts { + drafts[index] = promptDraft{Text: text, Images: nil} + } + + return drafts +} + +func promptDraftTexts(drafts []promptDraft) []string { + texts := make([]string, len(drafts)) + for index := range drafts { + texts[index] = drafts[index].Text + } + + return texts +} diff --git a/internal/terminal/panel_session_selection_internal_test.go b/internal/terminal/panel_session_selection_internal_test.go index 2edf5aa2..f0191781 100644 --- a/internal/terminal/panel_session_selection_internal_test.go +++ b/internal/terminal/panel_session_selection_internal_test.go @@ -67,7 +67,7 @@ func TestApplySessionSelectionAddsMessageAfterSuccessfulLoad(t *testing.T) { Role: database.RoleAssistant, Content: interruptTestPrompt, Provider: "", - Model: "", + Model: "", Parts: nil, }) if err != nil { t.Fatalf("append message: %v", err) diff --git a/internal/terminal/panel_test_helpers_internal_test.go b/internal/terminal/panel_test_helpers_internal_test.go index b5c4e82e..c57f43fc 100644 --- a/internal/terminal/panel_test_helpers_internal_test.go +++ b/internal/terminal/panel_test_helpers_internal_test.go @@ -46,7 +46,7 @@ func testEntryEntity() database.EntryEntity { Role: "", Content: "", Provider: "", - Model: "", + Model: "", Parts: nil, }, } } diff --git a/internal/terminal/panel_tree_internal_test.go b/internal/terminal/panel_tree_internal_test.go index a3e4d9b0..49a8fad2 100644 --- a/internal/terminal/panel_tree_internal_test.go +++ b/internal/terminal/panel_tree_internal_test.go @@ -85,7 +85,7 @@ func newTreePanelTestApp(ctx context.Context, t *testing.T) (app *App, userEntry Role: database.RoleUser, Content: interruptTestPrompt, Provider: "", - Model: "", + Model: "", Parts: nil, }) if err != nil { t.Fatalf("append user entry: %v", err) @@ -96,7 +96,7 @@ func newTreePanelTestApp(ctx context.Context, t *testing.T) (app *App, userEntry Role: database.RoleAssistant, Content: "world", Provider: "", - Model: "", + Model: "", Parts: nil, }) if err != nil { t.Fatalf("append assistant entry: %v", err) diff --git a/internal/terminal/prompt_cancel_internal_test.go b/internal/terminal/prompt_cancel_internal_test.go index 2678e45c..9b062d9c 100644 --- a/internal/terminal/prompt_cancel_internal_test.go +++ b/internal/terminal/prompt_cancel_internal_test.go @@ -23,7 +23,7 @@ func TestCancelActivePromptPreservesQueuedMessages(t *testing.T) { app.working = true app.addMessage(transcript.RoleUser, "prompt") app.appendStreamingBlock(transcript.RoleAssistant, "partial") - app.queuedMessages = []string{"follow up"} + app.queuedMessages = promptDrafts("follow up") app.activePrompt = newTestActivePrompt(func() { canceled = true }) app.activePrompt.Prompt = "prompt" @@ -36,7 +36,7 @@ func TestCancelActivePromptPreservesQueuedMessages(t *testing.T) { assert.Equal(t, "prompt", app.transcript.History[0].Content) require.Len(t, app.transcript.Streaming.Blocks, 1) assert.Equal(t, "partial", app.transcript.Streaming.Blocks[0].Content) - assert.Equal(t, []string{"follow up"}, app.queuedMessages) + assert.Equal(t, []string{"follow up"}, promptDraftTexts(app.queuedMessages)) assert.Equal(t, "canceling response...", app.statusMessage) } diff --git a/internal/terminal/prompt_history.go b/internal/terminal/prompt_history.go index 1d9a509c..e9d59494 100644 --- a/internal/terminal/prompt_history.go +++ b/internal/terminal/prompt_history.go @@ -1,6 +1,7 @@ package terminal import ( + "slices" "strings" "github.com/gdamore/tcell/v3" @@ -31,13 +32,14 @@ func (app *App) showPreviousPrompt() bool { if app.promptHistoryIndex == len(app.promptHistory) { app.promptHistoryDraft = app.composerBuffer.TextValue() + app.promptHistoryDraftImages = cloneImageAttachments(app.composerImages) } if app.promptHistoryIndex > 0 { app.promptHistoryIndex-- } - app.composerBuffer.SetText(app.promptHistory[app.promptHistoryIndex]) + app.restorePromptHistoryIndex(app.promptHistoryIndex) return true } @@ -49,25 +51,47 @@ func (app *App) showNextPrompt() bool { if app.promptHistoryIndex < len(app.promptHistory)-1 { app.promptHistoryIndex++ - app.composerBuffer.SetText(app.promptHistory[app.promptHistoryIndex]) + app.restorePromptHistoryIndex(app.promptHistoryIndex) return true } app.promptHistoryIndex = len(app.promptHistory) app.composerBuffer.SetText(app.promptHistoryDraft) + app.composerImages = cloneImageAttachments(app.promptHistoryDraftImages) app.promptHistoryDraft = "" + app.promptHistoryDraftImages = nil return true } +func (app *App) restorePromptHistoryIndex(index int) { + app.composerBuffer.SetText(app.promptHistory[index]) + + app.composerImages = nil + if index < len(app.promptHistoryImages) { + app.composerImages = cloneImageAttachments(app.promptHistoryImages[index]) + } +} + func (app *App) recordPromptHistory(text string) { - trimmed := strings.TrimSpace(text) - if trimmed == "" { + app.recordPromptDraftHistory(promptDraft{Text: text, Images: nil}) +} + +func (app *App) recordPromptDraftHistory(draft promptDraft) { + trimmed := strings.TrimSpace(draft.Text) + if trimmed == "" && len(draft.Images) == 0 { return } - if len(app.promptHistory) > 0 && app.promptHistory[len(app.promptHistory)-1] == trimmed { + last := len(app.promptHistory) - 1 + + var lastImages []imageAttachment + if last >= 0 && last < len(app.promptHistoryImages) { + lastImages = app.promptHistoryImages[last] + } + + if last >= 0 && app.promptHistory[last] == trimmed && attachmentSummariesEqual(lastImages, draft.Images) { app.resetPromptHistoryNavigation() return @@ -75,18 +99,38 @@ func (app *App) recordPromptHistory(text string) { if len(app.promptHistory) == promptHistoryLimit { app.promptHistory = append(app.promptHistory[:0], app.promptHistory[1:]...) + app.promptHistoryImages = append(app.promptHistoryImages[:0], app.promptHistoryImages[1:]...) } app.promptHistory = append(app.promptHistory, trimmed) + app.promptHistoryImages = append(app.promptHistoryImages, cloneImageAttachments(draft.Images)) app.resetPromptHistoryNavigation() } func (app *App) resetPromptHistory() { app.promptHistory = []string{} + app.promptHistoryImages = [][]imageAttachment{} app.resetPromptHistoryNavigation() } func (app *App) resetPromptHistoryNavigation() { app.promptHistoryIndex = len(app.promptHistory) app.promptHistoryDraft = "" + app.promptHistoryDraftImages = nil +} + +func attachmentSummariesEqual(left, right []imageAttachment) bool { + if len(left) != len(right) { + return false + } + + for index := range left { + if left[index].Name != right[index].Name || left[index].MIMEType != right[index].MIMEType || + left[index].Width != right[index].Width || left[index].Height != right[index].Height || + !slices.Equal(left[index].Data, right[index].Data) { + return false + } + } + + return true } diff --git a/internal/terminal/prompt_history_internal_test.go b/internal/terminal/prompt_history_internal_test.go index a03c004d..b6d29c24 100644 --- a/internal/terminal/prompt_history_internal_test.go +++ b/internal/terminal/prompt_history_internal_test.go @@ -27,6 +27,36 @@ func TestPromptHistoryNavigatesPromptsAndRestoresDraft(t *testing.T) { assertEditorText(t, app, "draft prompt") } +func TestPromptHistoryRestoresImagesAndCurrentDraft(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + historyImage := imageAttachment{ + Name: "history.png", MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1, + } + draftImage := imageAttachment{ + Name: "draft.png", MIMEType: clipboardImageMIME, Data: []byte{2}, Width: 2, Height: 2, + } + + app.recordPromptDraftHistory(promptDraft{Text: "with image", Images: []imageAttachment{historyImage}}) + app.composerBuffer.SetText("draft") + app.composerImages = []imageAttachment{draftImage} + + pressTerminalKey(t, app, tcell.KeyUp, "") + assertEditorText(t, app, "with image") + + if len(app.composerImages) != 1 || app.composerImages[0].Name != historyImage.Name { + t.Fatalf("history images = %#v", app.composerImages) + } + + pressTerminalKey(t, app, tcell.KeyDown, "") + assertEditorText(t, app, "draft") + + if len(app.composerImages) != 1 || app.composerImages[0].Name != draftImage.Name { + t.Fatalf("draft images = %#v", app.composerImages) + } +} + func TestPromptHistoryEditBecomesDraft(t *testing.T) { t.Parallel() diff --git a/internal/terminal/prompt_queue.go b/internal/terminal/prompt_queue.go index fc239f6b..c2b19f68 100644 --- a/internal/terminal/prompt_queue.go +++ b/internal/terminal/prompt_queue.go @@ -6,34 +6,51 @@ import ( ) func (app *App) queueFollowUp() { - text := strings.TrimSpace(app.composerBuffer.Clear()) - if text == "" { + draft := app.consumeDraft() + if draft.empty() { app.setStatus("no follow-up text to queue") return } - app.recordPromptHistory(text) - app.queueFollowUpText(text) -} + if !app.validateDraftModel(draft) { + app.restoreDraft(draft) + + return + } + + if strings.HasPrefix(draft.Text, "/") && len(draft.Images) > 0 { + app.restoreDraft(draft) + app.setStatus("slash commands do not accept image attachments") -func (app *App) queueFollowUpText(text string) { - app.queuePrompt(text, true) + return + } + + app.recordPromptDraftHistory(draft) + app.queueDraft(draft, true) } +func (app *App) queueFollowUpText(text string) { app.queuePrompt(text, true) } + +// queuePrompt keeps text-only internal workflow callers source compatible. func (app *App) queuePrompt(text string, visible bool) { - text = strings.TrimSpace(text) - if text == "" { + app.queueDraft(promptDraft{Text: strings.TrimSpace(text), Images: nil}, visible) +} + +func (app *App) queueDraft(draft promptDraft, visible bool) { + draft.Text = strings.TrimSpace(draft.Text) + if draft.empty() { return } + draft = clonePromptDraft(draft) if visible { - app.queuedMessages = append(app.queuedMessages, text) + app.queuedMessages = append(app.queuedMessages, draft) return } - app.hiddenQueuedMessages = append(app.hiddenQueuedMessages, text) + app.hiddenQueuedMessages = append(app.hiddenQueuedMessages, draft) } func (app *App) processQueuedPrompt(ctx context.Context) { @@ -42,9 +59,9 @@ func (app *App) processQueuedPrompt(ctx context.Context) { } if len(app.hiddenQueuedMessages) > 0 { - text := app.hiddenQueuedMessages[0] + draft := app.hiddenQueuedMessages[0] app.hiddenQueuedMessages = app.hiddenQueuedMessages[1:] - app.sendPromptHidden(ctx, text) + app.sendDraft(ctx, draft, false) return } @@ -53,28 +70,28 @@ func (app *App) processQueuedPrompt(ctx context.Context) { return } - text := app.queuedMessages[0] + draft := app.queuedMessages[0] app.queuedMessages = app.queuedMessages[1:] - app.sendPrompt(ctx, text) + app.sendDraft(ctx, draft, true) } -func (app *App) queuedCompactionPrompts() []string { +func (app *App) queuedCompactionPrompts() []promptDraft { if app.activeCompaction == nil || app.activeCompaction.QueuedStart >= len(app.queuedMessages) { return nil } - queued := append([]string(nil), app.queuedMessages[app.activeCompaction.QueuedStart:]...) + queued := clonePromptDrafts(app.queuedMessages[app.activeCompaction.QueuedStart:]) app.queuedMessages = app.queuedMessages[:app.activeCompaction.QueuedStart] return queued } -func (app *App) restoreCompactionQueuedPrompts(queued []string) { +func (app *App) restoreCompactionQueuedPrompts(queued []promptDraft) { if len(queued) == 0 { return } - app.queuedMessages = append(app.queuedMessages, queued...) + app.queuedMessages = append(app.queuedMessages, clonePromptDrafts(queued)...) app.dequeueFollowUp() } @@ -85,9 +102,9 @@ func (app *App) dequeueFollowUp() { return } - lastIndex := len(app.queuedMessages) - 1 + last := len(app.queuedMessages) - 1 app.resetPromptHistoryNavigation() - app.composerBuffer.SetText(app.queuedMessages[lastIndex]) - app.queuedMessages = app.queuedMessages[:lastIndex] + app.restoreDraft(app.queuedMessages[last]) + app.queuedMessages = app.queuedMessages[:last] app.setStatus("restored queued message") } diff --git a/internal/terminal/prompt_queue_internal_test.go b/internal/terminal/prompt_queue_internal_test.go index 06c48662..d974f08a 100644 --- a/internal/terminal/prompt_queue_internal_test.go +++ b/internal/terminal/prompt_queue_internal_test.go @@ -4,9 +4,17 @@ import ( "context" "slices" "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/omarluq/librecode/internal/model" ) -const testQueuedPromptText = "next prompt" +const ( + testQueuedImageName = "paste.png" + testQueuedPromptText = "next prompt" +) func TestQueueFollowUpText(t *testing.T) { t.Parallel() @@ -35,13 +43,38 @@ func TestQueueFollowUpText(t *testing.T) { t.Fatalf("statusMessage = %q, want %q", got, want) } - if got := app.queuedMessages; !slices.Equal(got, testCase.want) { + if got := app.queuedMessages; !slices.Equal(promptDraftTexts(got), testCase.want) { t.Fatalf("queuedMessages = %v, want %v", got, testCase.want) } }) } } +func TestQueueFollowUpRejectsImagesForTextOnlyModel(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + app.models = model.NewRegistry(&model.RegistryOptions{ + ConfigReader: nil, Auth: nil, ModelsPath: "", + BuiltIns: []model.Model{terminalCapabilityTestModel(app.currentProvider(), app.currentModel())}, + Discovery: disabledModelDiscovery(), + }) + app.composerImages = []imageAttachment{{ + Name: testQueuedImageName, MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1, + }} + app.composerBuffer.SetText("follow-up") + + app.queueFollowUp() + + if len(app.queuedMessages) != 0 { + t.Fatalf("queuedMessages length = %d, want 0", len(app.queuedMessages)) + } + + if len(app.composerImages) != 1 { + t.Fatalf("composerImages length = %d, want 1", len(app.composerImages)) + } +} + func TestQueueFollowUpRequiresText(t *testing.T) { t.Parallel() @@ -69,7 +102,7 @@ func TestQueueFollowUpRecordsAndClearsComposer(t *testing.T) { t.Fatalf("composer text = %q, want empty", got) } - if got, want := app.queuedMessages, []string{"follow me"}; !slices.Equal(got, want) { + if got, want := app.queuedMessages, []string{"follow me"}; !slices.Equal(promptDraftTexts(got), want) { t.Fatalf("queuedMessages = %v, want %v", got, want) } @@ -78,11 +111,57 @@ func TestQueueFollowUpRecordsAndClearsComposer(t *testing.T) { } } +func TestQueueAndDequeueFollowUpPreserveImages(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + image := imageAttachment{ + Name: testQueuedImageName, MIMEType: clipboardImageMIME, Data: []byte{1, 2}, Width: 1, Height: 1, + } + + app.composerBuffer.SetText("follow me") + app.composerImages = []imageAttachment{image} + + app.queueFollowUp() + require.Len(t, app.queuedMessages, 1) + require.Len(t, app.queuedMessages[0].Images, 1) + assert.Equal(t, image.Data, app.queuedMessages[0].Images[0].Data) + assert.True(t, app.composerDraftEmpty()) + + image.Data[0] = 9 + assert.Equal(t, byte(1), app.queuedMessages[0].Images[0].Data[0]) + + app.dequeueFollowUp() + assert.Equal(t, "follow me", app.composerBuffer.TextValue()) + require.Len(t, app.composerImages, 1) + assert.Equal(t, []byte{1, 2}, app.composerImages[0].Data) + assert.Empty(t, app.queuedMessages) +} + +func TestQueueAndDequeueImageOnlyFollowUp(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + app.composerImages = []imageAttachment{{ + Name: testQueuedImageName, MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1, + }} + + app.queueFollowUp() + require.Len(t, app.queuedMessages, 1) + assert.Empty(t, app.queuedMessages[0].Text) + assert.True(t, app.composerDraftEmpty()) + + app.dequeueFollowUp() + assert.Empty(t, app.composerBuffer.TextValue()) + require.Len(t, app.composerImages, 1) + assert.Equal(t, []byte{1}, app.composerImages[0].Data) +} + func TestDequeueFollowUpRestoresLastMessage(t *testing.T) { t.Parallel() app := newRenderTestApp(t) - app.queuedMessages = []string{"first", "second"} + app.queuedMessages = promptDrafts("first", "second") app.promptHistoryIndex = 1 app.dequeueFollowUp() @@ -91,7 +170,7 @@ func TestDequeueFollowUpRestoresLastMessage(t *testing.T) { t.Fatalf("composer text = %q, want %q", got, want) } - if got, want := app.queuedMessages, []string{"first"}; !slices.Equal(got, want) { + if got, want := app.queuedMessages, []string{"first"}; !slices.Equal(promptDraftTexts(got), want) { t.Fatalf("queuedMessages = %v, want %v", got, want) } @@ -154,7 +233,7 @@ func TestProcessQueuedPrompt(t *testing.T) { name: "busy leaves queue unchanged", setup: func(app *App) { app.working = true - app.queuedMessages = []string{firstQueuedPrompt} + app.queuedMessages = promptDrafts(firstQueuedPrompt) }, wantQueued: []string{firstQueuedPrompt}, wantWorking: true, @@ -167,7 +246,7 @@ func TestProcessQueuedPrompt(t *testing.T) { }, { name: sendFirstQueuedPrompt, - setup: func(app *App) { app.queuedMessages = []string{firstQueuedPrompt, secondQueuedPrompt} }, + setup: func(app *App) { app.queuedMessages = promptDrafts(firstQueuedPrompt, secondQueuedPrompt) }, wantQueued: []string{secondQueuedPrompt}, wantWorking: true, }, @@ -192,7 +271,7 @@ func TestProcessQueuedPrompt(t *testing.T) { _ = readPromptAsyncEvent(t, app) } - if !slices.Equal(app.queuedMessages, testCase.wantQueued) { + if !slices.Equal(promptDraftTexts(app.queuedMessages), testCase.wantQueued) { t.Fatalf("queuedMessages = %v, want %v", app.queuedMessages, testCase.wantQueued) } diff --git a/internal/terminal/prompt_response_internal_test.go b/internal/terminal/prompt_response_internal_test.go index 19134672..f86d245a 100644 --- a/internal/terminal/prompt_response_internal_test.go +++ b/internal/terminal/prompt_response_internal_test.go @@ -108,7 +108,7 @@ func TestApplyPromptResponseAddsAssistantAndProcessesQueue(t *testing.T) { client := newTerminalPromptClient(newTerminalCompletionResult("queued response"), nil) app := newPromptSendTestApp(t, client) app.activePrompt = newTestActivePrompt(nil) - app.queuedMessages = []string{asyncTestQueuedText} + app.queuedMessages = promptDrafts(asyncTestQueuedText) app.applyPromptResponse(context.Background(), newTestPromptResponse("assistant response"), app.activePrompt.ID) diff --git a/internal/terminal/prompt_send.go b/internal/terminal/prompt_send.go index 72345c1f..da37060c 100644 --- a/internal/terminal/prompt_send.go +++ b/internal/terminal/prompt_send.go @@ -10,16 +10,16 @@ import ( ) func (app *App) sendPrompt(ctx context.Context, text string) { - app.sendPromptWithVisibility(ctx, text, true) + app.sendDraft(ctx, promptDraft{Text: text, Images: nil}, true) } func (app *App) sendPromptHidden(ctx context.Context, text string) { - app.sendPromptWithVisibility(ctx, text, false) + app.sendDraft(ctx, promptDraft{Text: text, Images: nil}, false) } -func (app *App) sendPromptWithVisibility(ctx context.Context, text string, visible bool) { +func (app *App) sendDraft(ctx context.Context, draft promptDraft, visible bool) { if app.busy() { - app.queuePrompt(text, visible) + app.queueDraft(draft, visible) return } @@ -34,7 +34,8 @@ func (app *App) sendPromptWithVisibility(ctx context.Context, text string, visib ParentEntryID: parentEntryID, SessionID: app.sessionID, CWD: app.cwd, - Text: text, + Images: draft.assistantImages(), + Text: draft.Text, Name: "", ResumeLatest: false, HideUserPrompt: !visible, @@ -52,11 +53,14 @@ func (app *App) sendPromptWithVisibility(ctx context.Context, text string, visib ID: promptID, SessionID: app.sessionID, UserEntryID: "", - Prompt: text, + Prompt: draft.Text, + Images: cloneImageAttachments(draft.Images), Canceled: false, } if visible { - app.addMessage(transcript.RoleUser, text) + message := newChatMessage(transcript.RoleUser, draft.Text) + message.Attachments = summarizeAttachments(draft.Images) + app.appendMessage(message) } app.working = true diff --git a/internal/terminal/prompt_send_internal_test.go b/internal/terminal/prompt_send_internal_test.go index 2d6ca096..eb54c452 100644 --- a/internal/terminal/prompt_send_internal_test.go +++ b/internal/terminal/prompt_send_internal_test.go @@ -193,7 +193,7 @@ func assertSubmitCase( assert.Equal(t, testCase.wantMode, app.mode) assert.Equal(t, testCase.wantComposerText, app.composerBuffer.TextValue()) assert.Len(t, app.promptHistory, testCase.wantPromptHistory) - assertQueuedMessages(t, testCase.wantQueued, app.queuedMessages) + assertQueuedMessages(t, testCase.wantQueued, promptDraftTexts(app.queuedMessages)) assert.Equal(t, testCase.wantRequest, client.request != nil) } @@ -217,7 +217,7 @@ func TestSendPromptQueuesWhenWorking(t *testing.T) { app.sendPrompt(context.Background(), testQueuedPromptText) - assert.Equal(t, []string{testQueuedPromptText}, app.queuedMessages) + assert.Equal(t, []string{testQueuedPromptText}, promptDraftTexts(app.queuedMessages)) assert.Nil(t, app.activePrompt) } @@ -251,6 +251,33 @@ func TestSendPromptInitializesPromptState(t *testing.T) { assert.Equal(t, promptSendTestText, app.activePrompt.Prompt) } +func TestSubmitImageOnlyDraftSendsImagePart(t *testing.T) { + t.Parallel() + + client := newTerminalPromptClient(newTerminalCompletionResult("assistant response"), nil) + app := newPromptSendTestApp(t, client) + imageData := testPNG(t, 2, 3) + app.composerImages = []imageAttachment{{ + Name: "paste-1.png", MIMEType: clipboardImageMIME, Data: imageData, Width: 2, Height: 3, + }} + + shouldQuit, err := app.submit(t.Context()) + require.NoError(t, err) + assert.False(t, shouldQuit) + assert.True(t, app.composerDraftEmpty()) + require.NotNil(t, app.activePrompt) + require.Len(t, app.activePrompt.Images, 1) + assert.Equal(t, imageData, app.activePrompt.Images[0].Data) + + request := waitForPromptRequest(t, client) + require.NotEmpty(t, request.Messages) + userMessage := request.Messages[len(request.Messages)-1] + assert.Empty(t, userMessage.Content) + require.Len(t, userMessage.Parts, 1) + assert.Equal(t, database.MessagePartImage, userMessage.Parts[0].Type) + assert.Equal(t, imageData, userMessage.Parts[0].Data) +} + func TestRunPromptPostsDoneAndError(t *testing.T) { t.Parallel() @@ -288,6 +315,7 @@ func TestRunPromptPostsDoneAndError(t *testing.T) { ParentEntryID: nil, SessionID: "", CWD: app.cwd, + Images: nil, Text: promptSendTestText, Name: "", ResumeLatest: false, @@ -395,7 +423,7 @@ func promptSendTestModelDefinition() model.Model { Name: promptSendTestModel, API: "openai-completions", BaseURL: "https://example.invalid/v1", - Input: []model.InputMode{model.InputText}, + Input: []model.InputMode{model.InputText, model.InputImage}, Cost: model.Cost{Input: 0, Output: 0, CacheRead: 0, CacheWrite: 0}, ContextWindow: 1000, MaxTokens: 0, diff --git a/internal/terminal/prompt_submit.go b/internal/terminal/prompt_submit.go index 28cf8ed5..792d1abb 100644 --- a/internal/terminal/prompt_submit.go +++ b/internal/terminal/prompt_submit.go @@ -6,8 +6,12 @@ import ( ) func (app *App) submit(ctx context.Context) (bool, error) { - text := strings.TrimSpace(app.composerBuffer.TextValue()) - if text == "" { + draft := app.currentDraft() + if draft.empty() { + return false, nil + } + + if !app.draftCanSubmit(draft) { return false, nil } @@ -16,39 +20,62 @@ func (app *App) submit(ctx context.Context) (bool, error) { return false, err } - text = strings.TrimSpace(app.composerBuffer.Clear()) - if text == "" { + draft = app.currentDraft() + if !app.draftCanSubmit(draft) { + return false, nil + } + + draft = app.consumeDraft() + if draft.empty() { return false, nil } if app.compacting { - if strings.HasPrefix(text, "/") { - app.composerBuffer.SetText(text) - app.setStatus("wait for context compaction to finish") + return app.submitDuringCompaction(draft) + } - return false, nil - } + app.recordPromptDraftHistory(draft) - app.recordPromptHistory(text) - app.queueFollowUpText(text) - app.setStatus("queued prompt until context compaction finishes") + if strings.HasPrefix(draft.Text, "/") { + return app.submitCommand(ctx, draft.Text) + } + + if app.working { + app.queueDraft(draft, true) return false, nil } - app.recordPromptHistory(text) + app.sendDraft(ctx, draft, true) - if strings.HasPrefix(text, "/") { - return app.submitCommand(ctx, text) + return false, nil +} + +func (app *App) draftCanSubmit(draft promptDraft) bool { + if !app.validateDraftModel(draft) { + return false } - if app.working { - app.queueFollowUpText(text) + if strings.HasPrefix(draft.Text, "/") && len(draft.Images) > 0 { + app.setStatus("slash commands do not accept image attachments") + + return false + } + + return true +} + +func (app *App) submitDuringCompaction(draft promptDraft) (bool, error) { + if strings.HasPrefix(draft.Text, "/") { + app.restoreDraft(draft) + app.setStatus("wait for context compaction to finish") return false, nil } - app.sendPrompt(ctx, text) + app.recordPromptDraftHistory(draft) + app.queueDraft(draft, true) + app.setStatus("queued prompt until context compaction finishes") return false, nil } diff --git a/internal/terminal/render_composer.go b/internal/terminal/render_composer.go index afe4e31d..714d9701 100644 --- a/internal/terminal/render_composer.go +++ b/internal/terminal/render_composer.go @@ -43,10 +43,71 @@ func (app *App) drawComposerWindow(layout *extui.Layout) { } func (app *App) renderComposerEditor(width, bodyRows int) tui.TextAreaRender { - return app.composerBuffer.Render(width, bodyRows, tui.TextAreaStyles{ + chips := app.visibleAttachmentChipLines(width, max(0, bodyRows-1)) + textRows := max(1, bodyRows-len(chips)) + + rendered := app.composerBuffer.Render(width, textRows, tui.TextAreaStyles{ Border: app.theme.style(app.editorBorderColor()), Body: app.theme.style(colorText), }) + if len(chips) == 0 { + return rendered + } + + const afterTopBorder = 1 + + lines := make([]tui.Line, 0, len(rendered.Lines)+len(chips)) + lines = append(lines, rendered.Lines[:afterTopBorder]...) + lines = append(lines, chips...) + lines = append(lines, rendered.Lines[afterTopBorder:]...) + rendered.Lines = lines + rendered.CursorRow += len(chips) + + return rendered +} + +func (app *App) attachmentChipLines(width int) []tui.Line { + lines := make([]tui.Line, 0, len(app.composerImages)) + for _, item := range app.composerImages { + summary := attachmentSummary{ + Name: item.Name, MIMEType: item.MIMEType, Width: item.Width, + Height: item.Height, Size: len(item.Data), + } + text := " " + attachmentSummaryText(summary) + lines = append(lines, tui.NewLine(app.theme.style(colorDim), tui.Truncate(text, width))) + } + + return lines +} + +func (app *App) visibleAttachmentChipLines(width, limit int) []tui.Line { + chips := app.attachmentChipLines(width) + if limit <= 0 || len(chips) == 0 { + return nil + } + + if len(chips) <= limit { + return chips + } + + visible := max(0, limit-1) + lines := chips[:visible] + overflow := " … " + tui.Int(len(chips)-visible) + " more attachments" + + return append(lines, tui.NewLine(app.theme.style(colorDim), tui.Truncate(overflow, width))) +} + +func formatByteSize(size int) string { + const ( + kibibyte = 1024 + mebibyte = kibibyte * kibibyte + ) + + if size >= mebibyte { + return tui.Int((size+mebibyte-1)/mebibyte) + " MiB" + } + + return tui.Int((size+kibibyte-1)/kibibyte) + " KiB" } func (app *App) drawStatusWindow(layout *extui.Layout) { @@ -84,7 +145,7 @@ func (app *App) drawEditorAndFooter(width, height, _ int) { app.writeStyledLine(layout.footerStart+index, width, line) } - if app.transcriptListFocused() || app.agentTaskSummaryFocused() { + if len(layout.editor.Lines) == 0 || app.transcriptListFocused() || app.agentTaskSummaryFocused() { app.screen.HideCursor() return @@ -119,12 +180,15 @@ func (app *App) composerLayout(width, height int) composerLayout { if reserve > height { bodyRows := max(1, height-len(footerLines)-len(autocompleteLines)-composerBorderRows) editor = app.renderComposerEditor(width, bodyRows) - reserve = len(footerLines) + len(autocompleteLines) + len(editor.Lines) } + footerLines, autocompleteLines, editor = fitComposerRows( + height, footerLines, autocompleteLines, editor, + ) + reserve = len(footerLines) + len(autocompleteLines) + len(editor.Lines) startRow := max(0, height-reserve) editorStart := startRow + len(autocompleteLines) - footerStart := height - len(footerLines) + footerStart := editorStart + len(editor.Lines) return composerLayout{ editor: editor, @@ -137,6 +201,36 @@ func (app *App) composerLayout(width, height int) composerLayout { } } +func fitComposerRows( + height int, + footerLines, autocompleteLines []tui.Line, + editor tui.TextAreaRender, +) (footer, autocomplete []tui.Line, fittedEditor tui.TextAreaRender) { + height = max(0, height) + if len(footerLines) > height { + footerLines = footerLines[:height] + } + + remaining := height - len(footerLines) + if len(autocompleteLines) > remaining { + autocompleteLines = autocompleteLines[:remaining] + } + + remaining -= len(autocompleteLines) + if len(editor.Lines) > remaining { + editor.Lines = editor.Lines[:remaining] + } + + if len(editor.Lines) == 0 { + editor.CursorRow = 0 + editor.CursorCol = 0 + } else { + editor.CursorRow = min(editor.CursorRow, len(editor.Lines)-1) + } + + return footerLines, autocompleteLines, editor +} + func (app *App) editorBorderColor() colorToken { if strings.HasPrefix(strings.TrimSpace(app.composerBuffer.TextValue()), "!") { return colorBashMode diff --git a/internal/terminal/render_internal_test.go b/internal/terminal/render_internal_test.go index a2484b85..1e104b27 100644 --- a/internal/terminal/render_internal_test.go +++ b/internal/terminal/render_internal_test.go @@ -20,6 +20,31 @@ import ( "github.com/omarluq/librecode/internal/tui" ) +func TestComposerLayoutFitsShortTerminals(t *testing.T) { + t.Parallel() + + for height := 0; height <= 5; height++ { + t.Run(tui.Int(height), func(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + app.composerImages = []imageAttachment{ + {Name: "one.png", MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1}, + {Name: "two.png", MIMEType: clipboardImageMIME, Data: []byte{2}, Width: 1, Height: 1}, + {Name: "three.png", MIMEType: clipboardImageMIME, Data: []byte{3}, Width: 1, Height: 1}, + } + + layout := app.composerLayout(24, height) + assert.LessOrEqual(t, layout.reserve, height) + assert.GreaterOrEqual(t, layout.startRow, 0) + assert.GreaterOrEqual(t, layout.editorStart, 0) + assert.GreaterOrEqual(t, layout.footerStart, 0) + assert.LessOrEqual(t, layout.footerStart+len(layout.footerLines), height) + assert.LessOrEqual(t, layout.editorStart+len(layout.editor.Lines), height) + }) + } +} + func TestClearWindowRespectsWindowOrigin(t *testing.T) { t.Parallel() @@ -228,14 +253,14 @@ func TestRenderQueuedMessagesRendersHeadersAndBody(t *testing.T) { t.Parallel() tests := []struct { - queuedMessages []string + queuedMessages []promptDraft name string expectedLines []string width int }{ { name: "single queued message", - queuedMessages: []string{"first queued"}, + queuedMessages: promptDrafts("first queued"), width: 40, expectedLines: []string{ " queued follow-up 1 ", @@ -244,7 +269,7 @@ func TestRenderQueuedMessagesRendersHeadersAndBody(t *testing.T) { }, { name: "multiple queued messages", - queuedMessages: []string{"first queued", "second queued"}, + queuedMessages: promptDrafts("first queued", "second queued"), width: 40, expectedLines: []string{ " queued follow-up 1 ", @@ -1351,7 +1376,7 @@ func TestLoadInitialMessagesUsesTranscriptHistory(t *testing.T) { Role: database.RoleAssistant, Content: "loaded assistant", Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) diff --git a/internal/terminal/render_parity_internal_test.go b/internal/terminal/render_parity_internal_test.go index 052b7910..58184246 100644 --- a/internal/terminal/render_parity_internal_test.go +++ b/internal/terminal/render_parity_internal_test.go @@ -314,7 +314,7 @@ func appendRenderParityMessage( Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, }) requireNoError(t, err) diff --git a/internal/terminal/running_tools_internal_test.go b/internal/terminal/running_tools_internal_test.go index 7d2f7a32..bf68d1f6 100644 --- a/internal/terminal/running_tools_internal_test.go +++ b/internal/terminal/running_tools_internal_test.go @@ -143,7 +143,8 @@ func TestAgentTaskCompletionEventDrawsCollapsedExpandableToolResult(t *testing.T app := newRenderTestApp(t) app.working = true app.activePrompt = &activePromptState{ - Cancel: func() {}, ParentEntryID: nil, SessionID: "", UserEntryID: "", Prompt: "", ID: 1, Canceled: false, + Cancel: func() {}, ParentEntryID: nil, SessionID: "", UserEntryID: "", + Images: nil, Prompt: "", ID: 1, Canceled: false, } app.scrollOffset = 10 app.agentTasks = []database.AgentTaskEntity{testAgentTask(database.TaskRunning)} @@ -211,7 +212,8 @@ func TestAgentCompletionSurvivesPromptStreamingReset(t *testing.T) { app := newRenderTestApp(t) app.activePrompt = &activePromptState{ - Cancel: func() {}, ParentEntryID: nil, SessionID: "", UserEntryID: "", Prompt: "", ID: 1, Canceled: false, + Cancel: func() {}, ParentEntryID: nil, SessionID: "", UserEntryID: "", + Images: nil, Prompt: "", ID: 1, Canceled: false, } content := formatAgentCompletionForUI("Agent explore finished.\n\nreview complete") app.addAgentCompletionMessage(content) @@ -234,7 +236,8 @@ func TestAgentCompletionStaysLiveAcrossQueuedContinuation(t *testing.T) { app := newRenderTestApp(t) app.activePrompt = &activePromptState{ - Cancel: func() {}, ParentEntryID: nil, SessionID: "", UserEntryID: "", Prompt: "", ID: 1, Canceled: false, + Cancel: func() {}, ParentEntryID: nil, SessionID: "", UserEntryID: "", + Images: nil, Prompt: "", ID: 1, Canceled: false, } content := formatAgentCompletionForUI("Agent explore finished.\n\nreview complete") app.addAgentCompletionMessage(content) diff --git a/internal/terminal/session_commands_internal_test.go b/internal/terminal/session_commands_internal_test.go index 2a1ae970..e3cf1f14 100644 --- a/internal/terminal/session_commands_internal_test.go +++ b/internal/terminal/session_commands_internal_test.go @@ -117,7 +117,7 @@ func appendSessionMessage(t *testing.T, app *App, sessionID string, role databas Role: role, Content: content, Provider: "", - Model: "", + Model: "", Parts: nil, }) require.NoError(t, err) } diff --git a/internal/terminal/session_view.go b/internal/terminal/session_view.go index 5c5c2b59..201d1509 100644 --- a/internal/terminal/session_view.go +++ b/internal/terminal/session_view.go @@ -13,31 +13,34 @@ import ( // inspected. Prompt ownership remains on activePrompt and is intentionally not // part of a view. type sessionViewState struct { - lastEscape time.Time - pendingParentID *string - scopedEnabled map[string]bool - promptHistoryDraft string - streamingThinkingText string - streamingText string - statusMessage string - runningToolBlocks []runningToolBlock - queuedMessages []string - liveAgentCompletions []chatMessage - promptHistory []string - hiddenQueuedMessages []string - scopedOrder []string - settings sessionSettingsDocument - composerBuffer tui.TextArea - transcript transcriptState - tokenUsage model.TokenUsage - selection mouseSelection - transcriptList transcriptListSelection - streamedToolEvents int - promptHistoryIndex int - scrollOffset int - autocompleteSelection int - escapePresses int - autocompleteClosed bool + lastEscape time.Time + pendingParentID *string + scopedEnabled map[string]bool + promptHistoryDraft string + promptHistoryDraftImages []imageAttachment + streamingThinkingText string + streamingText string + statusMessage string + runningToolBlocks []runningToolBlock + queuedMessages []promptDraft + liveAgentCompletions []chatMessage + promptHistory []string + promptHistoryImages [][]imageAttachment + hiddenQueuedMessages []promptDraft + scopedOrder []string + settings sessionSettingsDocument + composerBuffer tui.TextArea + composerImages []imageAttachment + transcript transcriptState + tokenUsage model.TokenUsage + selection mouseSelection + transcriptList transcriptListSelection + streamedToolEvents int + promptHistoryIndex int + scrollOffset int + autocompleteSelection int + escapePresses int + autocompleteClosed bool } func (app *App) saveSessionView() { @@ -54,36 +57,44 @@ func (app *App) saveSessionView() { func (app *App) captureSessionView(clone bool) sessionViewState { view := sessionViewState{ - pendingParentID: app.pendingParentID, - transcript: app.transcript, - runningToolBlocks: app.runningToolBlocks, - liveAgentCompletions: app.liveAgentCompletions, - queuedMessages: app.queuedMessages, - hiddenQueuedMessages: app.hiddenQueuedMessages, - promptHistory: app.promptHistory, - promptHistoryDraft: app.promptHistoryDraft, - tokenUsage: app.tokenUsage, - composerBuffer: app.composerBuffer, - selection: app.selection, - transcriptList: app.transcriptList, - streamingText: app.streamingText, - streamingThinkingText: app.streamingThinkingText, - scopedEnabled: app.scopedEnabled, - scopedOrder: app.scopedOrder, - settings: app.currentSessionSettings(), - streamedToolEvents: app.streamedToolEvents, - promptHistoryIndex: app.promptHistoryIndex, - scrollOffset: app.scrollOffset, - statusMessage: app.statusMessage, - autocompleteSelection: app.autocompleteSelection, - autocompleteClosed: app.autocompleteClosed, - escapePresses: app.escapePresses, - lastEscape: app.lastEscape, + pendingParentID: app.pendingParentID, + transcript: app.transcript, + runningToolBlocks: app.runningToolBlocks, + liveAgentCompletions: app.liveAgentCompletions, + queuedMessages: app.queuedMessages, + hiddenQueuedMessages: app.hiddenQueuedMessages, + promptHistory: app.promptHistory, + promptHistoryImages: app.promptHistoryImages, + promptHistoryDraft: app.promptHistoryDraft, + promptHistoryDraftImages: app.promptHistoryDraftImages, + tokenUsage: app.tokenUsage, + composerBuffer: app.composerBuffer, + composerImages: app.composerImages, + selection: app.selection, + transcriptList: app.transcriptList, + streamingText: app.streamingText, + streamingThinkingText: app.streamingThinkingText, + scopedEnabled: app.scopedEnabled, + scopedOrder: app.scopedOrder, + settings: app.currentSessionSettings(), + streamedToolEvents: app.streamedToolEvents, + promptHistoryIndex: app.promptHistoryIndex, + scrollOffset: app.scrollOffset, + statusMessage: app.statusMessage, + autocompleteSelection: app.autocompleteSelection, + autocompleteClosed: app.autocompleteClosed, + escapePresses: app.escapePresses, + lastEscape: app.lastEscape, } if clone { view.pendingParentID = cloneStringPtr(view.pendingParentID) view.tokenUsage = cloneTerminalUsage(view.tokenUsage) view.composerBuffer = cloneComposerBuffer(view.composerBuffer) + view.composerImages = cloneImageAttachments(view.composerImages) + view.promptHistoryImages = cloneImageAttachmentGroups(view.promptHistoryImages) + view.promptHistoryDraftImages = cloneImageAttachments(view.promptHistoryDraftImages) + view.queuedMessages = clonePromptDrafts(view.queuedMessages) + view.hiddenQueuedMessages = clonePromptDrafts(view.hiddenQueuedMessages) view.scopedEnabled = maps.Clone(view.scopedEnabled) view.scopedOrder = slices.Clone(view.scopedOrder) } @@ -111,9 +122,12 @@ func (app *App) applySessionView(sessionID string, view *sessionViewState, clone app.queuedMessages = view.queuedMessages app.hiddenQueuedMessages = view.hiddenQueuedMessages app.promptHistory = view.promptHistory + app.promptHistoryImages = view.promptHistoryImages app.promptHistoryDraft = view.promptHistoryDraft + app.promptHistoryDraftImages = view.promptHistoryDraftImages app.tokenUsage = view.tokenUsage app.composerBuffer = view.composerBuffer + app.composerImages = view.composerImages app.selection = view.selection app.transcriptList = view.transcriptList app.streamingText = view.streamingText @@ -134,6 +148,11 @@ func (app *App) applySessionView(sessionID string, view *sessionViewState, clone app.pendingParentID = cloneStringPtr(app.pendingParentID) app.tokenUsage = cloneTerminalUsage(app.tokenUsage) app.composerBuffer = cloneComposerBuffer(app.composerBuffer) + app.composerImages = cloneImageAttachments(app.composerImages) + app.promptHistoryImages = cloneImageAttachmentGroups(app.promptHistoryImages) + app.promptHistoryDraftImages = cloneImageAttachments(app.promptHistoryDraftImages) + app.queuedMessages = clonePromptDrafts(app.queuedMessages) + app.hiddenQueuedMessages = clonePromptDrafts(app.hiddenQueuedMessages) app.scopedEnabled = maps.Clone(app.scopedEnabled) app.scopedOrder = slices.Clone(app.scopedOrder) } diff --git a/internal/terminal/workflow_summary_internal_test.go b/internal/terminal/workflow_summary_internal_test.go index 2eb8a44f..d78595ac 100644 --- a/internal/terminal/workflow_summary_internal_test.go +++ b/internal/terminal/workflow_summary_internal_test.go @@ -130,7 +130,7 @@ func TestWorkflowFailureIsPushedIntoCompletedTurn(t *testing.T) { assert.Contains(t, app.liveAgentCompletions[0].Content, "failed-run") assert.Contains(t, app.liveAgentCompletions[0].Content, "compile failed") require.Len(t, app.hiddenQueuedMessages, 1) - assert.Contains(t, app.hiddenQueuedMessages[0], "background workflow failed") + assert.Contains(t, app.hiddenQueuedMessages[0].Text, "background workflow failed") app.refreshActiveWorkflows(t.Context()) assert.Len(t, app.liveAgentCompletions, 1, "failure must only be delivered once") @@ -226,7 +226,7 @@ func TestWorkflowFailureNotificationFallbacksAndBusyBranches(t *testing.T) { assert.Contains(t, app.liveAgentCompletions[0].Content, toolDisplayWorkflow) assert.Contains(t, app.liveAgentCompletions[0].Content, "No error detail was returned.") require.Len(t, app.hiddenQueuedMessages, 1) - assert.Contains(t, app.hiddenQueuedMessages[0], "background workflow failed") + assert.Contains(t, app.hiddenQueuedMessages[0].Text, "background workflow failed") }) } } From 5fbff3d6df2e38f185d0299b8d214a6b90d9b04e Mon Sep 17 00:00:00 2001 From: Omar Alani Date: Mon, 3 Aug 2026 23:34:11 -0500 Subject: [PATCH 2/6] fix(images): harden multipart prompt handling --- go.mod | 2 +- .../lifecyclepayload/lifecyclepayload_test.go | 44 +++++--- .../assistant/llm_conversion_internal_test.go | 97 ++++++++++++----- internal/assistant/prompt_images.go | 22 +++- .../assistant/prompt_images_internal_test.go | 75 +++++++++++++ internal/assistant/runtime.go | 4 +- internal/assistant/runtime_model.go | 16 ++- internal/assistant/runtime_slash.go | 2 +- .../assistant/runtime_slash_internal_test.go | 42 ++++++- .../assistant/test_constants_internal_test.go | 5 +- internal/assistant/testdata/prompt.webp | Bin 0 -> 4296 bytes .../00014_index_image_message_parts.sql | 7 ++ internal/database/migrations_test.go | 19 +++- internal/database/session_entry_repository.go | 29 ++++- .../database/session_message_parts_test.go | 103 +++++++++++++++++- .../database/session_message_repository.go | 58 +++++++++- .../sqlite_contention_internal_test.go | 2 +- internal/database/validation.go | 3 +- internal/provider/anthropic.go | 2 +- internal/provider/image_content.go | 11 ++ .../provider/image_content_internal_test.go | 43 ++++++++ internal/provider/messages.go | 6 +- internal/provider/openai_chat.go | 2 +- internal/terminal/agent_tasks.go | 8 +- .../agent_tasks_behavior_internal_test.go | 39 +++++++ .../terminal/attachments_internal_test.go | 38 ++++++- .../compact_commands_internal_test.go | 24 +++- .../extension_events_internal_test.go | 10 +- .../terminal/prompt_cancel_internal_test.go | 8 +- internal/terminal/session_view.go | 44 ++++---- .../terminal/session_view_internal_test.go | 19 ++++ 31 files changed, 667 insertions(+), 117 deletions(-) create mode 100644 internal/assistant/testdata/prompt.webp create mode 100644 internal/database/migrations/00014_index_image_message_parts.sql diff --git a/go.mod b/go.mod index ed804458..d9444e06 100644 --- a/go.mod +++ b/go.mod @@ -48,6 +48,7 @@ require ( github.com/yuin/goldmark v1.8.5 github.com/yuin/gopher-lua v1.1.2 golang.design/x/clipboard v0.8.0 + golang.org/x/image v0.41.0 golang.org/x/net v0.57.0 golang.org/x/text v0.40.0 gopkg.in/yaml.v3 v3.0.1 @@ -107,7 +108,6 @@ require ( golang.design/x/x11 v0.2.0 // indirect golang.org/x/exp v0.0.0-20260718201538-764159d718ef // indirect golang.org/x/exp/shiny v0.0.0-20250606033433-dcc06ee1d476 // indirect - golang.org/x/image v0.41.0 // indirect golang.org/x/mobile v0.0.0-20250606033058-a2a15c67f36f // indirect golang.org/x/sync v0.22.0 // indirect golang.org/x/sys v0.47.0 // indirect diff --git a/internal/assistant/lifecyclepayload/lifecyclepayload_test.go b/internal/assistant/lifecyclepayload/lifecyclepayload_test.go index 1c145636..fe8ed087 100644 --- a/internal/assistant/lifecyclepayload/lifecyclepayload_test.go +++ b/internal/assistant/lifecyclepayload/lifecyclepayload_test.go @@ -28,25 +28,37 @@ const ( func TestPromptAndTurnPayloads(t *testing.T) { t.Parallel() - t.Run("prompt", func(t *testing.T) { + t.Run("prompt attachments", func(t *testing.T) { t.Parallel() parentID := lifecycleTestParentID - prompt := lifecyclepayload.Prompt(&lifecyclepayload.PromptRequest{ - ParentEntryID: &parentID, - Attachments: nil, - CWD: "/work", - Name: "agent", - SessionID: lifecycleTestSessionID, - Text: "hello", - ResumeLatest: true, - }) - assert.Equal(t, "/work", prompt[lifecyclepayload.CWDKey]) - assert.Equal(t, "agent", prompt[lifecyclepayload.ToolNameKey]) - assert.Equal(t, parentID, prompt[lifecyclepayload.ParentEntryIDKey]) - resumeLatest, ok := prompt["resume_latest"].(bool) - require.True(t, ok) - assert.True(t, resumeLatest) + + tests := []struct { + name string + attachments []map[string]any + wantCount int + }{ + {name: "none", attachments: nil, wantCount: 0}, + {name: "one", attachments: []map[string]any{{"name": "screen.webp", "size": 442}}, wantCount: 1}, + } + for _, testCase := range tests { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + prompt := lifecyclepayload.Prompt(&lifecyclepayload.PromptRequest{ + ParentEntryID: &parentID, Attachments: testCase.attachments, CWD: "/work", + Name: "agent", SessionID: lifecycleTestSessionID, Text: "hello", ResumeLatest: true, + }) + assert.Equal(t, "/work", prompt[lifecyclepayload.CWDKey]) + assert.Equal(t, "agent", prompt[lifecyclepayload.ToolNameKey]) + assert.Equal(t, parentID, prompt[lifecyclepayload.ParentEntryIDKey]) + assert.Equal(t, testCase.wantCount, prompt["attachment_count"]) + assert.Equal(t, testCase.attachments, prompt["attachments"]) + resumeLatest, ok := prompt["resume_latest"].(bool) + require.True(t, ok) + assert.True(t, resumeLatest) + }) + } }) t.Run("nil prompt", func(t *testing.T) { diff --git a/internal/assistant/llm_conversion_internal_test.go b/internal/assistant/llm_conversion_internal_test.go index 58d22139..69295dd1 100644 --- a/internal/assistant/llm_conversion_internal_test.go +++ b/internal/assistant/llm_conversion_internal_test.go @@ -27,7 +27,7 @@ func TestLLMRequestFromCompletionRequestConvertsAssistantState(t *testing.T) { ToolRegistry: registry, ExecuteTools: nil, SessionID: "session-1", - SystemPrompt: jsonSystemRole, + SystemPrompt: "system instructions", ThinkingLevel: "high", CWD: t.TempDir(), Auth: model.RequestAuth{ @@ -67,10 +67,12 @@ func TestLLMRequestFromCompletionRequestConvertsAssistantState(t *testing.T) { converted := llmRequestFromCompletionRequest(request) assert.Equal(t, "session-1", converted.SessionID) - assert.Equal(t, expectedSystemRole, converted.SystemPrompt) + assert.Equal(t, "system instructions", converted.SystemPrompt) assert.Equal(t, "high", converted.ThinkingLevel) assert.Equal(t, "openai", converted.Model.Provider) assert.Equal(t, "gpt-test", converted.Model.ID) + assert.Equal(t, apiOpenAIResponses, converted.Model.API) + assert.Equal(t, "https://example.test", converted.Model.BaseURL) assert.Equal(t, "secret", converted.Auth.APIKey) assert.Equal(t, "value", converted.Auth.Headers["x-test"]) assert.Len(t, converted.Messages, 4) @@ -87,36 +89,77 @@ func TestLLMRequestFromCompletionRequestConvertsAssistantState(t *testing.T) { assert.Equal(t, "yes", request.Model.Compat["compat"]) } -func TestLLMMessageFromDatabasePreservesOrderedMultipartImageOnly(t *testing.T) { +func TestLLMMessageFromDatabaseConvertsContentParts(t *testing.T) { t.Parallel() data := []byte{0, 1, 2, 3} - message, converted := llmMessageFromDatabase(&database.MessageEntity{ - Timestamp: time.Time{}, Role: database.RoleUser, Content: "", Provider: "", Model: "", - Parts: []database.MessagePartEntity{ - {Text: "inspect image", MIMEType: "", Name: "", Type: database.MessagePartText, - Data: nil, Width: 0, Height: 0}, - {Text: "", MIMEType: imageMIMEPNG, Name: "screen.png", Type: database.MessagePartImage, - Data: data, Width: 10, Height: 20}, + tests := []struct { + name string + wantText string + message database.MessageEntity + wantTypes []llm.PartType + }{ + { + name: "ordered multipart", + message: database.MessageEntity{ + Timestamp: time.Time{}, Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{ + {Text: "inspect image", MIMEType: "", Name: "", Type: database.MessagePartText, + Data: nil, Width: 0, Height: 0}, + {Text: "", MIMEType: imageMIMEPNG, Name: "screen.png", Type: database.MessagePartImage, + Data: data, Width: 10, Height: 20}, + }, + }, + wantTypes: []llm.PartType{llm.PartText, llm.PartImage}, + wantText: "inspect image", }, - }) + { + name: "image only", + message: database.MessageEntity{ + Timestamp: time.Time{}, Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{{ + Text: "", MIMEType: imageMIMEPNG, Name: "", Type: database.MessagePartImage, + Data: data, Width: 1, Height: 1, + }}, + }, + wantTypes: []llm.PartType{llm.PartImage}, + wantText: "", + }, + { + name: "legacy content fallback", + message: database.MessageEntity{ + Timestamp: time.Time{}, Role: database.RoleAssistant, Content: "answer", + Provider: "", Model: "", Parts: nil, + }, + wantTypes: []llm.PartType{llm.PartText}, + wantText: "answer", + }, + } - require.True(t, converted) - require.Len(t, message.Content, 2) - assert.Equal(t, llm.PartText, message.Content[0].Type) - assert.Equal(t, llm.PartImage, message.Content[1].Type) - assert.Equal(t, base64.StdEncoding.EncodeToString(data), message.Content[1].Data) - assert.Equal(t, "screen.png", message.Content[1].Metadata["name"]) - - imageOnly, imageOnlyConverted := llmMessageFromDatabase(&database.MessageEntity{ - Timestamp: time.Time{}, Role: database.RoleUser, Content: "", Provider: "", Model: "", - Parts: []database.MessagePartEntity{{ - Text: "", MIMEType: imageMIMEPNG, Name: "", Type: database.MessagePartImage, - Data: data, Width: 1, Height: 1, - }}, - }) - require.True(t, imageOnlyConverted) - assert.Len(t, imageOnly.Content, 1) + for _, testCase := range tests { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + message, converted := llmMessageFromDatabase(&testCase.message) + require.True(t, converted) + require.Len(t, message.Content, len(testCase.wantTypes)) + + for index := range testCase.wantTypes { + assert.Equal(t, testCase.wantTypes[index], message.Content[index].Type) + } + + if testCase.wantText != "" { + assert.Equal(t, testCase.wantText, message.Content[0].Text) + } + + if message.Content[len(message.Content)-1].Type == llm.PartImage { + imagePart := message.Content[len(message.Content)-1] + assert.Equal(t, base64.StdEncoding.EncodeToString(data), imagePart.Data) + assert.Equal(t, imageMIMEPNG, imagePart.MIMEType) + assert.NotEmpty(t, imagePart.Metadata) + } + }) + } } func TestLLMRequestFromCompletionRequestNilAndDisabledTools(t *testing.T) { diff --git a/internal/assistant/prompt_images.go b/internal/assistant/prompt_images.go index 3f379391..47bf369c 100644 --- a/internal/assistant/prompt_images.go +++ b/internal/assistant/prompt_images.go @@ -10,6 +10,7 @@ import ( "strings" "github.com/samber/oops" + _ "golang.org/x/image/webp" // Register WebP for DecodeConfig. "github.com/omarluq/librecode/internal/database" "github.com/omarluq/librecode/internal/model" @@ -22,6 +23,7 @@ const ( maxPromptImageTotal = 20 << 20 maxPromptImagePixels = 40_000_000 imageMIMEPNG = "image/png" + imageMIMEWebP = "image/webp" ) func (runtime *Runtime) preparePromptRequest(request *PromptRequest) (*PromptRequest, error) { @@ -170,6 +172,8 @@ func imageMIMEType(format string) string { return "image/jpeg" case "png": return imageMIMEPNG + case "webp": + return imageMIMEWebP default: return "" } @@ -207,11 +211,19 @@ func validateSelectedModelHasImageInput(selected *model.Model, source string) er } func messagesContainImages(messages []database.MessageEntity) bool { - for messageIndex := range messages { - for partIndex := range messages[messageIndex].Parts { - if messages[messageIndex].Parts[partIndex].Type == database.MessagePartImage { - return true - } + for index := range messages { + if messageContainsImages(&messages[index]) { + return true + } + } + + return false +} + +func messageContainsImages(message *database.MessageEntity) bool { + for index := range message.Parts { + if message.Parts[index].Type == database.MessagePartImage { + return true } } diff --git a/internal/assistant/prompt_images_internal_test.go b/internal/assistant/prompt_images_internal_test.go index f8ec35ee..fb4c352e 100644 --- a/internal/assistant/prompt_images_internal_test.go +++ b/internal/assistant/prompt_images_internal_test.go @@ -4,6 +4,7 @@ import ( "bytes" "image" "image/png" + "os" "testing" "time" @@ -15,6 +16,8 @@ import ( "github.com/omarluq/librecode/internal/model" ) +const testWebPName = "screen.webp" + func testPNG(t *testing.T, width, height int) []byte { t.Helper() @@ -24,6 +27,31 @@ func testPNG(t *testing.T, width, height int) []byte { return output.Bytes() } +func testWebP(t *testing.T) []byte { + t.Helper() + + data, err := os.ReadFile("testdata/prompt.webp") + require.NoError(t, err) + + return data +} + +func TestValidatePromptImagesSupportsWebP(t *testing.T) { + t.Parallel() + + attachment := ImageAttachment{ + Name: testWebPName, MIMEType: imageMIMEWebP, Data: testWebP(t), Width: 75, Height: 100, + } + require.NoError(t, validatePromptImages([]ImageAttachment{attachment})) + + attachment.MIMEType = imageMIMEPNG + err := validatePromptImages([]ImageAttachment{attachment}) + require.Error(t, err) + coded, ok := oops.AsOops(err) + require.True(t, ok) + assert.Equal(t, "invalid_image_mime", coded.Code()) +} + func TestValidatePromptImagesAndCloneBoundary(t *testing.T) { t.Parallel() @@ -40,6 +68,53 @@ func TestValidatePromptImagesAndCloneBoundary(t *testing.T) { assert.NotEqual(t, cloned.Images[0].Data[0], request.Images[0].Data[0]) } +func TestLifecyclePromptRequestAttachments(t *testing.T) { + t.Parallel() + + tests := []struct { + request *PromptRequest + name string + wantName string + wantCount int + }{ + {name: "nil request", request: nil, wantName: "", wantCount: 0}, + { + name: "no images", + request: &PromptRequest{ + OnEvent: nil, OnRetry: nil, OnUserEntry: nil, ParentEntryID: nil, + SessionID: "", CWD: "", Text: "", Images: nil, Name: "", + ResumeLatest: false, HideUserPrompt: false, + }, + wantName: "", wantCount: 0, + }, + { + name: "image metadata", + request: &PromptRequest{ + OnEvent: nil, OnRetry: nil, OnUserEntry: nil, ParentEntryID: nil, + SessionID: "", CWD: "", Text: "", Name: "", ResumeLatest: false, HideUserPrompt: false, + Images: []ImageAttachment{{ + Name: testWebPName, MIMEType: imageMIMEWebP, Data: []byte{1, 2}, Width: 3, Height: 4, + }}, + }, + wantCount: 1, wantName: testWebPName, + }, + } + + for _, testCase := range tests { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + payload := lifecyclePromptRequest(testCase.request) + assert.Len(t, payload.Attachments, testCase.wantCount) + + if testCase.wantCount > 0 { + assert.Equal(t, testCase.wantName, payload.Attachments[0][executeNameKey]) + assert.Equal(t, 2, payload.Attachments[0]["size"]) + } + }) + } +} + func TestValidateSelectedModelImageInput(t *testing.T) { t.Parallel() diff --git a/internal/assistant/runtime.go b/internal/assistant/runtime.go index 46685581..981a5fbc 100644 --- a/internal/assistant/runtime.go +++ b/internal/assistant/runtime.go @@ -520,7 +520,9 @@ func (runtime *Runtime) ModelRegistry() *model.Registry { } func splitSlashCommand(prompt string) (name, args string) { - trimmedPrompt := strings.TrimSpace(strings.TrimPrefix(prompt, slashPrefix)) + trimmedPrompt := strings.TrimPrefix(strings.TrimSpace(prompt), slashPrefix) + + trimmedPrompt = strings.TrimSpace(trimmedPrompt) if trimmedPrompt == "" { return "", "" } diff --git a/internal/assistant/runtime_model.go b/internal/assistant/runtime_model.go index 19c4eeaf..c8eecb3e 100644 --- a/internal/assistant/runtime_model.go +++ b/internal/assistant/runtime_model.go @@ -46,7 +46,12 @@ func (runtime *Runtime) respond( err error, ) { if strings.HasPrefix(strings.TrimSpace(prompt), slashPrefix) { - slashResponse, slashToolEvents, slashErr := runtime.respondToSlashCommand(ctx, cwd, prompt, onEvent) + slashResponse, slashToolEvents, slashErr := runtime.respondToSlashCommand( + ctx, + cwd, + strings.TrimSpace(prompt), + onEvent, + ) return &responseBundle{ Text: slashResponse, @@ -119,6 +124,8 @@ func (runtime *Runtime) modelResponse( } if contextHasImages { + // Keep this early check: historical images are known before auth and context + // construction, so an incompatible model should fail without doing either. imageErr := validateSelectedModelHasImageInput(&selectedModel, "conversation_history") if imageErr != nil { return nil, imageErr @@ -429,12 +436,13 @@ func (runtime *Runtime) promptContextContainsImages( return false, nil } - contextEntity, err := runtime.modelContextEntityFrom(ctx, sessionID, lineage.activeParentEntryID) + hasImages, err := runtime.sessions.ContextHasImageParts(ctx, sessionID, lineage.activeParentEntryID) if err != nil { - return false, err + return false, oops.In("assistant").Code("load_image_context"). + Wrapf(err, "query session context for images") } - return messagesContainImages(contextEntity.Messages), nil + return hasImages, nil } func (runtime *Runtime) cacheKey(sessionID, prompt string) string { diff --git a/internal/assistant/runtime_slash.go b/internal/assistant/runtime_slash.go index 3486f0fb..0cb0bf57 100644 --- a/internal/assistant/runtime_slash.go +++ b/internal/assistant/runtime_slash.go @@ -25,7 +25,7 @@ func (runtime *Runtime) respondToSlashCommand( prompt string, onEvent func(StreamEvent), ) (response string, toolEvents []ToolEvent, err error) { - commandName, commandArgs := splitSlashCommand(prompt) + commandName, commandArgs := splitSlashCommand(strings.TrimSpace(prompt)) if commandName == "" { return "", nil, oops.In("assistant").Code("empty_slash_command").Errorf("empty slash command") } diff --git a/internal/assistant/runtime_slash_internal_test.go b/internal/assistant/runtime_slash_internal_test.go index c38f0444..1afaed9a 100644 --- a/internal/assistant/runtime_slash_internal_test.go +++ b/internal/assistant/runtime_slash_internal_test.go @@ -18,9 +18,40 @@ import ( const ( runtimeSlashHelloPrompt = "/hello world" + runtimeSlashHelloReply = "hi world" runtimeSlashMissing = "missing" ) +func TestSplitSlashCommandRemovesOnePrefix(t *testing.T) { + t.Parallel() + + const ( + slashTestArgs = "fast" + slashTestModel = "model" + ) + + tests := []struct { + name string + prompt string + wantName string + wantArgs string + }{ + {name: "single prefix", prompt: "/model fast", wantName: slashTestModel, wantArgs: slashTestArgs}, + {name: "double prefix", prompt: "//model fast", wantName: "/model", wantArgs: slashTestArgs}, + {name: "surrounding whitespace", prompt: " /model fast ", wantName: slashTestModel, wantArgs: slashTestArgs}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + name, args := splitSlashCommand(test.prompt) + assert.Equal(t, test.wantName, name) + assert.Equal(t, test.wantArgs, args) + }) + } +} + func TestRespondToSlashCommand(t *testing.T) { t.Parallel() @@ -39,10 +70,17 @@ func TestRespondToSlashCommand(t *testing.T) { wantErrText: "empty slash command", }, { - extensions: slashCommandExtensions{err: nil, response: "hi world"}, + extensions: slashCommandExtensions{err: nil, response: runtimeSlashHelloReply}, prompt: runtimeSlashHelloPrompt, name: "extension command", - wantResponse: "hi world", + wantResponse: runtimeSlashHelloReply, + wantErrText: "", + }, + { + extensions: slashCommandExtensions{err: nil, response: runtimeSlashHelloReply}, + prompt: " " + runtimeSlashHelloPrompt + " ", + name: "extension command with surrounding whitespace", + wantResponse: runtimeSlashHelloReply, wantErrText: "", }, { diff --git a/internal/assistant/test_constants_internal_test.go b/internal/assistant/test_constants_internal_test.go index 34bab52d..dca7bfd2 100644 --- a/internal/assistant/test_constants_internal_test.go +++ b/internal/assistant/test_constants_internal_test.go @@ -1,6 +1,3 @@ package assistant -const ( - expectedReadToolName = string(jsonReadToolName) - expectedSystemRole = string(jsonSystemRole) -) +const expectedReadToolName = string(jsonReadToolName) diff --git a/internal/assistant/testdata/prompt.webp b/internal/assistant/testdata/prompt.webp new file mode 100644 index 0000000000000000000000000000000000000000..fe5018d2764fb9a28864cb430739a0c42bc9ff64 GIT binary patch literal 4296 zcmaKvXE%S4YF1fygKqn98uL^lKx%!zLF5~3&a5GhI! zEjUPY5#=Q6A>w#P&U?M*`S3jNx~{d?Z~yQAz4lu7+8_2dzow_x#RP!0u9m5lsj>|X z06>=9(Ezvra7|m=v=ByGO=x?sCg|xA~iKO}hC;<;Z02*-4*(JbN z)6~@P%>AF+!32PuMaHDg=K9Z`|Ekivx&@F|1VNG=C>LMxpjR=FTFTMWODpsbxGJ(ozy7-c09! zb17?*!z5{1NkEx3p;U`>Uy(Igm}$y=B6q{xUfy-eh+pY@$s72*o79HE5gcVN8)F+` zszn#Nm=ezOSj=)Y1EC=Bqca)3lH~u22*=+qcAC0i%rScZOVk1Xc}*P@*nfZ|@Y1WsOwSOnpze`DUuexiIfin3$o6Dl;Vl0+B)=Xt`Fj%wIQYt~RZ`@rZ87 zkovhY1X6Oe4hpwFmd&M6dULQAr^WMeNU5veB#GpN#;`q8-=2t;o*~Vw6|(a1rT0BA z$Mg&{{*+Dh7=`Lk5_o z;?#F~{hn3BEFVJ^dB)3Z;V>@|xV&ywg-%_lkQZ^6O~fK;YkI$GY>5gs(AWD_)~P4% zk;UWb*}Yl2t4E{aZ80LW_J0MpdyW?R4{_yi`Caa6(Tl|tEr*ht)?}RrhJ{xnbZ2Vk z1pn2NLd&|!N4e`Zm;$q|{PA} z#uH)^2n6XHL)}J4V@<(F-FF+@-kN6)GN|~NBWUqeS&A%05`S&iS{WtvK}^M#kKCP- z!}`1u5e5)-HqeHA$#2eHDIfABKDMHlM4ymhzk21+XJ`W0BJfhR4ltp+cn zk4bpY=|2Yg^vw8!8Mh9(C|(Argx zr+VdqDxx#~Er)!w+tZ%%n4gXks!AF{7gd-kulUo14XnKnzOKM|fWGH76#dAa^Xm-i zJY!O3ZoI=ZfvZ?)kv7&sM7fP3w5>Ckg+{fh@2hFrqp@-}P49%Of{?0v^{Kf8m`_}~ z+Tw)wH#0F|+(gCR5Ss}G(>>QQzr;M!l8_XgA$0BsUB#M0eSf{s$fiotrLr+*PC6bX zQobBgFj#~u1=TghA~d@hlA*wU{))LI8yo_Qo7F*cc+{$oWZiU^yPsZ%ZH(pTyt|5x zC?W7(NW?wYqm&u693xqYS(S`j^u`hE3uy1Wvj`QViAqv5dvfq5O)q5NiQH29>D{#! z{Wu3p?X7ml6W6h#fkD<0!1)N3`|4xWqyChO{QkTh!7;R~VRUW%V;hfqRQa24iP6Hn zLU-83O^A{A@{h(9&L4z-wGIb;sVp{aLfl}RB^v1Q?Ab+g+k1V#1-XEJzcq#%r{53m zYp}~wKZ-DDXpoilMQ~T{3`h#*iv|vBUkixk{Vpe7-$I&sM<3`gVnnoUkX@&2u02p$ zD9l?|-ya~T$Piw9QFDA=hi8|QW>EcC#0LjX4`QUo8q_cu!t_<#0e?JA1sHH61_rf= zruC57o;Y@zET`@ZxM|TCHMOiwDYKOnnTq!G`Gxu9{oWwE9O`b}Y&SuGE`2wH(v8lK zzZBVYWwIbhdV9&0-EKmO7^2VPQT@R__UYdC`J&kOzqh^Uhk$r0IdJFPEQDwsOy#tA znMaTf4ojUrDlx^^PoHD`DXJ*R8&Up%U5Ek3FQEE>68}br#dl&+7T9m7^iNTwbS?`( zn=JLCoU8VX;qEsGKbi(5;-$KaYiDZ&s&WmqBW1gxOfs|_sH!h5OegrFZuMkXFwR3* z3)epUJu}ROTeVJdS9OYl(Ov71IJxf>%@U&P$H^g3NhJc(Q_7Cumqkzoqd|5TIkK?X zK}l==3!diEvsx^YTS@l@H9oYFzihPmdFlbal*VX>8}`h2Q_g_VF6Uyv4oFhcKjFAQ z*|upi)w1nUVXoHPdE!LR1I4FYSHeK&J&Wbmqh8qmC^mT1(U*Jhq+l@NwSk-&a!An~ z$mB@{&C>IjjJK^ke#pKEq4s~vA_C4aVskH}mQX&m(Vy4`Cdxmr@Aa3@jv;|vHAC)x z?3xO4+(dgR7r=iK4%5}}AK$e@K)W2z{ZRCh#kW7R+I;BmJiX^RVaI8UdTNxV_Xfga zugxSfba&ymZ}(QB1)gHZ=kW%OP)6>rk))z+q9T2zMoV0jRcE-@O$|Nco}yYbKUn06 z79B7m@_G3W@-iCc$la!jvi!sLzF3hRF21&7r0~Gp?N$J?xi5xK=-_m8qn+I&f42+i zQr<|HPK|-rw9%NYap(t&(j)agonsM^UVGKUwDeq9bE02*f*PP*f=QCef_T2>!1VIj z?^85mLMi#~-$Gvo4ry#X>_Xoe~_*8YR` zI|O*M{YmNzEru?`Jh*cW$sE;4IWnXWbw^wAOGWWs_%?0MFm^ynd-_PM~}_VG%( zZ!u87s{he;&{TPY5>mHg+`3rI>TvH;0bkdFEqKGu;YGc6f*v=gp<$okO=0R{FWXWBW4%q z?RwTANO%WxjmoYDzox_2S zD``Q3T~-=NEAvO*PA50yB99_UMi23cONKXB2XAyRsN%WGVu&H_%5R^b`8iQT ze_hK?w>lvtZ*l!m%iM4zKkCel>{$+J42=Ky{Z=B3QXAF%Lp{W$+Rn71=?<*dduatv z&BaqhLV$T`rnHN6>f-^LPXQ`5Oxm}e*ZQ|wS&dAFQK~d|pZyajt~Jg>%;5}lN6`h< z$i57KR%)}U?w`1X=7b2j?q#{lMMMrH?SvWv^*&Mp5U;)DrDLZ7ME9ttR*y&d7ZKRE z=J}pLZ`}sE<9#AL;XqOyU#~fKPFjVw^i#T?|Db;vA!}vCnC1wlNf;ygJbi6ZSDY(~ z{nWcxRl+1f2}$_yjzyaQiPJ2G-P%*;X0G`cPLA7cm%w*xZH;DHEl!0-Trdyg*Rm72)Nv z(v?_b8uA)#o*Wf%GQQc-H8G`KlMbu7>d!3ZkX2yIuwQkwQo%OT{DjmA4VW^tdB)Aa z{=~1qP#hQ^kMPr5L6BvSy2hLcDq!)dvBS8>&HS)Q9VlHOuB$s!v%44u1Z9`cS@>Kj z^cyn_wvBw-h4vX=_~tifwP_gt&@69fYkVE43NTM*UHSGVx#yD+k*}vVYu!uHZW*8R zmW@`*9!gTNu>ikGr{W4H%!f0dEA zPHc`3N|Je7u|4zZ#vJGx;xFp@6m1Qs;4J0c=A}QL{ld^1uYh5QdiDu{Qg+iJ2nVPh z{lv%BsE13XA}R0`B~~yb`U$7DW2+mopQ;GokMxu4cIXbFLp$wYxJ;Vi>a{j?4p9l_8}%IrlQzgmN4+3 znS!&DvY?V1d)wbtN>gdqm<3-Kc$@IVQs$Z%RT!B4KugWYE?mPR{@U!zCUdGIP0;M? zQLfnGO6O5dY`N#l=a0v;z zwi1`BQdV8I+Ba}9r4#u;M|d|{HjfQRTAI3hySJ05wBO1+Bz!J`4Fok%6#8dwxhL6F zexBSC{$x&BTbQswss)lw{?J|oenof5abTF(sh5oE5Gea|YsK%?!OGp*&ex5f%{)A7 z)H`l7UHW#XcVlC)f`;#RhG)EEa?dt0<8})K-Oq2ItF+#I(?rtA=cl len("image/") + strings.HasPrefix(value, "image/") && len(value) > len("image/") && + !strings.Contains(value, "*") } func validateTaskEntity(entity *TaskEntity) error { diff --git a/internal/provider/anthropic.go b/internal/provider/anthropic.go index a5700d86..4a9d2f02 100644 --- a/internal/provider/anthropic.go +++ b/internal/provider/anthropic.go @@ -460,7 +460,7 @@ func anthropicMessages(messages []llm.Message) ([]map[string]any, error) { content = anthropicUserContent(message) } - if content == "" { + if emptyMessageContent(content) { continue } diff --git a/internal/provider/image_content.go b/internal/provider/image_content.go index 7b071727..153828fb 100644 --- a/internal/provider/image_content.go +++ b/internal/provider/image_content.go @@ -128,6 +128,17 @@ func openAIChatUserContent(message llm.Message) any { return blocks } +func emptyMessageContent(content any) bool { + switch value := content.(type) { + case string: + return value == "" + case []map[string]any: + return len(value) == 0 + default: + return false + } +} + func anthropicUserContent(message llm.Message) any { hasImage := false for index := range message.Content { diff --git a/internal/provider/image_content_internal_test.go b/internal/provider/image_content_internal_test.go index 28396a51..adf5b24e 100644 --- a/internal/provider/image_content_internal_test.go +++ b/internal/provider/image_content_internal_test.go @@ -112,6 +112,26 @@ func TestMultipartProviderMappingsPreserveImageOnlyOrder(t *testing.T) { assert.Len(t, anthropicContent, 2) } +func TestMultipartProviderMappingsSkipEmptyUserMessages(t *testing.T) { + t.Parallel() + + empty := llm.Message{Metadata: nil, Role: llm.RoleUser, Content: []llm.Part{}} + + responses, err := openAIResponseInput([]llm.Message{empty}) + require.NoError(t, err) + assert.Empty(t, responses) + + request := emptyCompletionRequest() + setTestRequestMessages(request, []llm.Message{empty}) + chat, err := openAIChatMessages(request) + require.NoError(t, err) + assert.Empty(t, chat) + + anthropic, err := anthropicMessages([]llm.Message{empty}) + require.NoError(t, err) + assert.Empty(t, anthropic) +} + func TestMultipartProviderMappingsPreserveTextOnlyPayloads(t *testing.T) { t.Parallel() @@ -162,6 +182,29 @@ func TestMultipartProviderMappingsRejectMalformedAndUnsupportedRole(t *testing.T } } +func TestProvidersSkipEmptyStructuredUserContent(t *testing.T) { + t.Parallel() + + message := llm.Message{Metadata: nil, Role: llm.RoleUser, Content: []llm.Part{{ + Metadata: nil, ToolCall: nil, ToolResult: nil, Type: llm.PartFile, + Text: "", MIMEType: "", Data: "", + }}} + + responses, err := openAIResponseInput([]llm.Message{message}) + require.NoError(t, err) + assert.Empty(t, responses) + + request := emptyCompletionRequest() + setTestRequestMessages(request, []llm.Message{message}) + chat, err := openAIChatMessages(request) + require.NoError(t, err) + assert.Empty(t, chat) + + anthropic, err := anthropicMessages([]llm.Message{message}) + require.NoError(t, err) + assert.Empty(t, anthropic) +} + func TestCodexCompactionRejectsAssistantImagesBeforeFlattening(t *testing.T) { t.Parallel() diff --git a/internal/provider/messages.go b/internal/provider/messages.go index 000841af..b35fcaff 100644 --- a/internal/provider/messages.go +++ b/internal/provider/messages.go @@ -22,12 +22,10 @@ func openAIResponseInput(messages []llm.Message) ([]any, error) { var content any = messageText(message) if message.Role == llm.RoleUser { - if blocks := openAIResponseUserContent(message); len(blocks) > 0 { - content = blocks - } + content = openAIResponseUserContent(message) } - if content == "" { + if emptyMessageContent(content) { continue } diff --git a/internal/provider/openai_chat.go b/internal/provider/openai_chat.go index 837009b0..a9562e61 100644 --- a/internal/provider/openai_chat.go +++ b/internal/provider/openai_chat.go @@ -212,7 +212,7 @@ func openAIChatMessages(request *CompletionRequest) ([]map[string]any, error) { content = openAIChatUserContent(message) } - if content == "" { + if emptyMessageContent(content) { continue } diff --git a/internal/terminal/agent_tasks.go b/internal/terminal/agent_tasks.go index af2dbe96..ceefb1ae 100644 --- a/internal/terminal/agent_tasks.go +++ b/internal/terminal/agent_tasks.go @@ -1562,10 +1562,10 @@ func (app *App) switchToAgentTaskSession( app.agentTaskSessionStack = sessionStack if app.restoreSessionView(sessionID) { - promptHistory := app.promptHistory - promptHistoryImages := app.promptHistoryImages + promptHistory := slices.Clone(app.promptHistory) + promptHistoryImages := cloneImageAttachmentGroups(app.promptHistoryImages) promptHistoryDraft := app.promptHistoryDraft - promptHistoryDraftImages := app.promptHistoryDraftImages + promptHistoryDraftImages := cloneImageAttachments(app.promptHistoryDraftImages) promptHistoryIndex := app.promptHistoryIndex app.transcript.History = nil app.transcript.LineCache.reset() @@ -1710,7 +1710,7 @@ func (app *App) appendMissingSessionMessages(messages []database.SessionMessageE CreatedAt: message.CreatedAt, Role: role, Content: message.Content, - Attachments: nil, + Attachments: databaseAttachmentSummaries(message.Parts), }) appended = true diff --git a/internal/terminal/agent_tasks_behavior_internal_test.go b/internal/terminal/agent_tasks_behavior_internal_test.go index ea130082..7a9d8b2b 100644 --- a/internal/terminal/agent_tasks_behavior_internal_test.go +++ b/internal/terminal/agent_tasks_behavior_internal_test.go @@ -4,8 +4,10 @@ import ( "context" "database/sql" "errors" + "fmt" "os" "path/filepath" + "slices" "strings" "sync" "testing" @@ -752,6 +754,43 @@ func TestRevisitAgentTaskSessionRefreshesDurableTranscript(t *testing.T) { assert.Contains(t, app.transcript.History[0].Content, "new durable child message") } +func TestRevisitAgentTaskPreservesFullPromptHistoryDuringRefresh(t *testing.T) { + t.Parallel() + + fixture, _, app := newAgentTaskSessionTestApp(t, database.TaskSucceeded) + history := make([]string, promptHistoryLimit) + historyImages := make([][]imageAttachment, promptHistoryLimit) + + for index := range history { + history[index] = fmt.Sprintf("prompt %d", index) + historyImages[index] = []imageAttachment{{ + Name: fmt.Sprintf("image-%d.png", index), MIMEType: clipboardImageMIME, Data: []byte{byte(index)}, + Width: 1, Height: 1, + }} + } + + require.NoError(t, app.inspectAgentTask(t.Context(), behaviorTaskID)) + app.promptHistory = slices.Clone(history) + app.promptHistoryImages = cloneImageAttachmentGroups(historyImages) + require.NoError(t, app.leaveAgentTaskSession(t.Context())) + + _, err := fixture.sessions.AppendMessage(t.Context(), fixture.child.ID, nil, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "durable refresh prompt", + Provider: "", Model: "", Parts: nil, + }) + require.NoError(t, err) + + require.NoError(t, app.inspectAgentTask(t.Context(), behaviorTaskID)) + assert.Equal(t, history, app.promptHistory) + require.Len(t, app.promptHistoryImages, promptHistoryLimit) + + for index := range historyImages { + require.Len(t, app.promptHistoryImages[index], 1) + assert.Equal(t, historyImages[index][0].Name, app.promptHistoryImages[index][0].Name) + assert.Equal(t, historyImages[index][0].Data, app.promptHistoryImages[index][0].Data) + } +} + func TestRevisitRunningAgentTaskPreservesTransientState(t *testing.T) { t.Parallel() diff --git a/internal/terminal/attachments_internal_test.go b/internal/terminal/attachments_internal_test.go index 895905d0..a5f9b9e5 100644 --- a/internal/terminal/attachments_internal_test.go +++ b/internal/terminal/attachments_internal_test.go @@ -7,11 +7,13 @@ import ( "image/png" "strings" "testing" + "time" "github.com/gdamore/tcell/v3" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/omarluq/librecode/internal/database" "github.com/omarluq/librecode/internal/model" ) @@ -113,15 +115,23 @@ func TestBracketedPasteInsertsLiterally(t *testing.T) { func TestPromptDraftAndSessionViewCloneImageBytes(t *testing.T) { t.Parallel() app := newRenderTestApp(t) - app.sessionID = "session" + app.sessionID = benchmarkDisplayedSession app.composerImages = []imageAttachment{{Name: "image", MIMEType: "", Data: []byte{1, 2}, Width: 0, Height: 0}} + app.promptHistory = []string{"history"} + app.promptHistoryImages = [][]imageAttachment{{{ + Name: "history-image", MIMEType: "", Data: []byte{4}, Width: 0, Height: 0, + }}} queuedImage := imageAttachment{Name: "", MIMEType: "", Data: []byte{3}, Width: 0, Height: 0} app.queuedMessages = []promptDraft{{Text: "queued", Images: []imageAttachment{queuedImage}}} app.saveSessionView() app.composerImages[0].Data[0] = 9 + app.promptHistory[0] = "changed" + app.promptHistoryImages[0][0].Data[0] = 9 app.queuedMessages[0].Images[0].Data[0] = 9 - view := app.sessionViews["session"] + view := app.sessionViews[benchmarkDisplayedSession] assert.Equal(t, byte(1), view.composerImages[0].Data[0]) + assert.Equal(t, "history", view.promptHistory[0]) + assert.Equal(t, byte(4), view.promptHistoryImages[0][0].Data[0]) assert.Equal(t, byte(3), view.queuedMessages[0].Images[0].Data[0]) } @@ -219,6 +229,30 @@ func TestReadOnlyInspectionBlocksImageAndBracketedPaste(t *testing.T) { assert.Equal(t, readOnlyAgentInspectionStatus, app.statusMessage) } +func TestAppendSessionMessagesRestoresAttachmentSummaries(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + app.appendSessionMessages([]database.SessionMessageEntity{{ + CreatedAt: time.Time{}, ID: "", SessionID: "", EntryID: "", Sender: "", + Role: database.RoleUser, Content: "", Provider: "", Model: "", + Parts: []database.MessagePartEntity{{ + Type: database.MessagePartImage, Text: "", Name: testImageAttachmentName, + MIMEType: clipboardImageMIME, Data: []byte("image"), Width: 10, Height: 20, + }}, + }}) + + require.Len(t, app.transcript.History, 1) + summaries := app.transcript.History[0].Attachments + require.NotNil(t, summaries) + require.Len(t, *summaries, 1) + assert.Equal(t, testImageAttachmentName, (*summaries)[0].Name) + assert.Equal(t, clipboardImageMIME, (*summaries)[0].MIMEType) + assert.Equal(t, 10, (*summaries)[0].Width) + assert.Equal(t, 20, (*summaries)[0].Height) + assert.Equal(t, len("image"), (*summaries)[0].Size) +} + func TestAttachmentRenderingContainsMetadataNotData(t *testing.T) { t.Parallel() app := newRenderTestApp(t) diff --git a/internal/terminal/compact_commands_internal_test.go b/internal/terminal/compact_commands_internal_test.go index 64e00c5d..e4da5ecb 100644 --- a/internal/terminal/compact_commands_internal_test.go +++ b/internal/terminal/compact_commands_internal_test.go @@ -179,9 +179,12 @@ func TestHandleCompactDoneStartsQueuedPrompt(t *testing.T) { client := newTerminalPromptClient(newTerminalCompletionResult("queued response"), nil) app := newPromptSendTestApp(t, client) + image := imageAttachment{ + Name: testQueuedImageName, MIMEType: clipboardImageMIME, Data: testPNG(t, 2, 3), Width: 2, Height: 3, + } app.compacting = true app.activeCompaction = &activeCompactionState{Cancel: func() {}, ID: 9, QueuedStart: 0} - app.queuedMessages = promptDrafts("queued after compact") + app.queuedMessages = []promptDraft{{Text: "queued after compact", Images: []imageAttachment{image}}} app.applyCompactDone(context.Background(), &asyncEvent{ Response: nil, ToolCallEvent: nil, ToolEvent: nil, Usage: &model.TokenUsage{ @@ -196,7 +199,15 @@ func TestHandleCompactDoneStartsQueuedPrompt(t *testing.T) { }) request := waitForPromptRequest(t, client) - assert.Equal(t, "queued after compact", request.Messages[len(request.Messages)-1].Content) + message := request.Messages[len(request.Messages)-1] + assert.Equal(t, "queued after compact", message.Content) + require.Len(t, message.Parts, 2) + assert.Equal(t, database.MessagePartImage, message.Parts[1].Type) + assert.Equal(t, image.Name, message.Parts[1].Name) + assert.Equal(t, image.MIMEType, message.Parts[1].MIMEType) + assert.Equal(t, image.Data, message.Parts[1].Data) + assert.Equal(t, image.Width, message.Parts[1].Width) + assert.Equal(t, image.Height, message.Parts[1].Height) assert.Empty(t, app.queuedMessages) assert.Equal(t, 10_000, app.tokenUsage.ContextTokens) } @@ -524,7 +535,13 @@ func TestCompactErrorRestoresQueuedPrompt(t *testing.T) { t.Parallel() app := newRenderTestApp(t) - app.queuedMessages = promptDrafts("preexisting", "during compaction") + image := imageAttachment{ + Name: testQueuedImageName, MIMEType: clipboardImageMIME, Data: []byte{1, 2}, Width: 2, Height: 3, + } + app.queuedMessages = []promptDraft{ + {Text: "preexisting", Images: nil}, + {Text: "during compaction", Images: []imageAttachment{image}}, + } app.compacting = true app.activeCompaction = &activeCompactionState{Cancel: func() {}, ID: 9, QueuedStart: 1} @@ -534,6 +551,7 @@ func TestCompactErrorRestoresQueuedPrompt(t *testing.T) { }) assert.Equal(t, "during compaction", app.composerBuffer.TextValue()) + assert.Equal(t, []imageAttachment{image}, app.composerImages) assert.Equal(t, []string{"preexisting"}, promptDraftTexts(app.queuedMessages)) } diff --git a/internal/terminal/extension_events_internal_test.go b/internal/terminal/extension_events_internal_test.go index f5d40af0..d716b8d0 100644 --- a/internal/terminal/extension_events_internal_test.go +++ b/internal/terminal/extension_events_internal_test.go @@ -89,9 +89,11 @@ end) `) app.working = true app.composerBuffer.SetText("original") - app.composerImages = []imageAttachment{{ - Name: testImageAttachmentName, MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 1, Height: 1, - }} + + image := imageAttachment{ + Name: testImageAttachmentName, MIMEType: clipboardImageMIME, Data: []byte{1}, Width: 2, Height: 3, + } + app.composerImages = []imageAttachment{image} shouldQuit, err := app.submit(t.Context()) require.NoError(t, err) @@ -99,7 +101,7 @@ end) require.Len(t, app.queuedMessages, 1) assert.Equal(t, "mutated by extension", app.queuedMessages[0].Text) require.Len(t, app.queuedMessages[0].Images, 1) - assert.Equal(t, []byte{1}, app.queuedMessages[0].Images[0].Data) + assert.Equal(t, image, app.queuedMessages[0].Images[0]) } func TestExtensionPromptSubmitRevalidatesImageDraft(t *testing.T) { diff --git a/internal/terminal/prompt_cancel_internal_test.go b/internal/terminal/prompt_cancel_internal_test.go index 9b062d9c..cf8e9053 100644 --- a/internal/terminal/prompt_cancel_internal_test.go +++ b/internal/terminal/prompt_cancel_internal_test.go @@ -23,7 +23,11 @@ func TestCancelActivePromptPreservesQueuedMessages(t *testing.T) { app.working = true app.addMessage(transcript.RoleUser, "prompt") app.appendStreamingBlock(transcript.RoleAssistant, "partial") - app.queuedMessages = promptDrafts("follow up") + + image := imageAttachment{ + Name: testQueuedImageName, MIMEType: clipboardImageMIME, Data: []byte{1, 2}, Width: 2, Height: 3, + } + app.queuedMessages = []promptDraft{{Text: "follow up", Images: []imageAttachment{image}}} app.activePrompt = newTestActivePrompt(func() { canceled = true }) app.activePrompt.Prompt = "prompt" @@ -37,6 +41,8 @@ func TestCancelActivePromptPreservesQueuedMessages(t *testing.T) { require.Len(t, app.transcript.Streaming.Blocks, 1) assert.Equal(t, "partial", app.transcript.Streaming.Blocks[0].Content) assert.Equal(t, []string{"follow up"}, promptDraftTexts(app.queuedMessages)) + require.Len(t, app.queuedMessages[0].Images, 1) + assert.Equal(t, image, app.queuedMessages[0].Images[0]) assert.Equal(t, "canceling response...", app.statusMessage) } diff --git a/internal/terminal/session_view.go b/internal/terminal/session_view.go index 201d1509..859284e3 100644 --- a/internal/terminal/session_view.go +++ b/internal/terminal/session_view.go @@ -87,21 +87,26 @@ func (app *App) captureSessionView(clone bool) sessionViewState { lastEscape: app.lastEscape, } if clone { - view.pendingParentID = cloneStringPtr(view.pendingParentID) - view.tokenUsage = cloneTerminalUsage(view.tokenUsage) - view.composerBuffer = cloneComposerBuffer(view.composerBuffer) - view.composerImages = cloneImageAttachments(view.composerImages) - view.promptHistoryImages = cloneImageAttachmentGroups(view.promptHistoryImages) - view.promptHistoryDraftImages = cloneImageAttachments(view.promptHistoryDraftImages) - view.queuedMessages = clonePromptDrafts(view.queuedMessages) - view.hiddenQueuedMessages = clonePromptDrafts(view.hiddenQueuedMessages) - view.scopedEnabled = maps.Clone(view.scopedEnabled) - view.scopedOrder = slices.Clone(view.scopedOrder) + cloneSessionViewState(&view) } return view } +func cloneSessionViewState(view *sessionViewState) { + view.pendingParentID = cloneStringPtr(view.pendingParentID) + view.tokenUsage = cloneTerminalUsage(view.tokenUsage) + view.composerBuffer = cloneComposerBuffer(view.composerBuffer) + view.composerImages = cloneImageAttachments(view.composerImages) + view.promptHistory = slices.Clone(view.promptHistory) + view.promptHistoryImages = cloneImageAttachmentGroups(view.promptHistoryImages) + view.promptHistoryDraftImages = cloneImageAttachments(view.promptHistoryDraftImages) + view.queuedMessages = clonePromptDrafts(view.queuedMessages) + view.hiddenQueuedMessages = clonePromptDrafts(view.hiddenQueuedMessages) + view.scopedEnabled = maps.Clone(view.scopedEnabled) + view.scopedOrder = slices.Clone(view.scopedOrder) +} + func (app *App) restoreSessionView(sessionID string) bool { view, found := app.sessionViews[sessionID] if !found { @@ -114,6 +119,12 @@ func (app *App) restoreSessionView(sessionID string) bool { } func (app *App) applySessionView(sessionID string, view *sessionViewState, clone bool) { + if clone { + cloned := *view + cloneSessionViewState(&cloned) + view = &cloned + } + app.sessionID = sessionID app.pendingParentID = view.pendingParentID app.transcript = view.transcript @@ -143,19 +154,6 @@ func (app *App) applySessionView(sessionID string, view *sessionViewState, clone app.escapePresses = view.escapePresses app.lastEscape = view.lastEscape app.applySessionSettings(&view.settings) - - if clone { - app.pendingParentID = cloneStringPtr(app.pendingParentID) - app.tokenUsage = cloneTerminalUsage(app.tokenUsage) - app.composerBuffer = cloneComposerBuffer(app.composerBuffer) - app.composerImages = cloneImageAttachments(app.composerImages) - app.promptHistoryImages = cloneImageAttachmentGroups(app.promptHistoryImages) - app.promptHistoryDraftImages = cloneImageAttachments(app.promptHistoryDraftImages) - app.queuedMessages = clonePromptDrafts(app.queuedMessages) - app.hiddenQueuedMessages = clonePromptDrafts(app.hiddenQueuedMessages) - app.scopedEnabled = maps.Clone(app.scopedEnabled) - app.scopedOrder = slices.Clone(app.scopedOrder) - } } func (app *App) inspectingWhilePromptRuns() bool { diff --git a/internal/terminal/session_view_internal_test.go b/internal/terminal/session_view_internal_test.go index 75549529..082c3a39 100644 --- a/internal/terminal/session_view_internal_test.go +++ b/internal/terminal/session_view_internal_test.go @@ -55,6 +55,25 @@ func TestWithSessionViewRejectsEventWhenOwnerViewIsMissing(t *testing.T) { assert.Empty(t, app.transcript.History) } +func TestSessionViewSaveAndRestoreClonePromptHistory(t *testing.T) { + t.Parallel() + + app := newRenderTestApp(t) + + const savedPrompt = "saved prompt" + + app.sessionID = benchmarkDisplayedSession + app.promptHistory = []string{savedPrompt} + app.saveSessionView() + + app.promptHistory[0] = "mutated active prompt" + assert.Equal(t, []string{savedPrompt}, app.sessionViews[benchmarkDisplayedSession].promptHistory) + + require.True(t, app.restoreSessionView(benchmarkDisplayedSession)) + app.promptHistory[0] = "changed restored prompt" + assert.Equal(t, []string{savedPrompt}, app.sessionViews[benchmarkDisplayedSession].promptHistory) +} + func TestPromptEventReportsMissingOwnerView(t *testing.T) { t.Parallel() From 9591e18ac628208d0fd479b197f316d804a121a6 Mon Sep 17 00:00:00 2001 From: Omar Alani Date: Tue, 4 Aug 2026 10:23:32 -0500 Subject: [PATCH 3/6] fix(images): address review quality findings --- internal/assistant/runtime_model.go | 68 ++++---- internal/assistant/runtime_persist.go | 20 +-- .../00013_add_session_message_parts.sql | 10 +- .../00014_index_image_message_parts.sql | 2 +- internal/database/migrations_test.go | 89 ++++++++++- internal/database/session_entry_repository.go | 41 ++--- .../database/session_message_parts_test.go | 146 +++++++++++++----- internal/provider/image_content.go | 44 ++---- .../compact_commands_internal_test.go | 2 + 9 files changed, 279 insertions(+), 143 deletions(-) diff --git a/internal/assistant/runtime_model.go b/internal/assistant/runtime_model.go index c8eecb3e..4d06f936 100644 --- a/internal/assistant/runtime_model.go +++ b/internal/assistant/runtime_model.go @@ -21,6 +21,17 @@ type promptLineage struct { activeParentEntryID string } +type responseInput struct { + lineage *promptLineage + onEvent func(StreamEvent) + onRetry RetryEventHandler + sessionID string + cwd string + prompt string + hasPromptImages bool + contextHasImages bool +} + func newPromptLineage(userEntryID string) *promptLineage { return &promptLineage{activeParentEntryID: userEntryID} } @@ -33,24 +44,18 @@ func (lineage *promptLineage) adopt(entry *database.EntryEntity) { func (runtime *Runtime) respond( ctx context.Context, - sessionID string, - lineage *promptLineage, - cwd string, - prompt string, - hasImages bool, - onEvent func(StreamEvent), - onRetry RetryEventHandler, + input *responseInput, ) ( bundle *responseBundle, cached bool, err error, ) { - if strings.HasPrefix(strings.TrimSpace(prompt), slashPrefix) { + if strings.HasPrefix(strings.TrimSpace(input.prompt), slashPrefix) { slashResponse, slashToolEvents, slashErr := runtime.respondToSlashCommand( ctx, - cwd, - strings.TrimSpace(prompt), - onEvent, + input.cwd, + strings.TrimSpace(input.prompt), + input.onEvent, ) return &responseBundle{ @@ -62,19 +67,18 @@ func (runtime *Runtime) respond( }, false, slashErr } - cacheKey := runtime.cacheKey(sessionID, prompt) + cacheKey := runtime.cacheKey(input.sessionID, input.prompt) - contextHasImages, contextErr := runtime.promptContextContainsImages(ctx, sessionID, lineage) + contextHasImages, contextErr := runtime.promptContextContainsImages(ctx, input.sessionID, input.lineage) if contextErr != nil { return nil, false, contextErr } - if hasImages || contextHasImages { + if input.hasPromptImages || contextHasImages { // Image bytes deliberately stay out of cache keys; prompts with image-bearing // context always execute against their durable multipart history. - imageBundle, modelErr := runtime.modelResponse( - ctx, sessionID, lineage, cwd, prompt, contextHasImages, onEvent, onRetry, - ) + input.contextHasImages = contextHasImages + imageBundle, modelErr := runtime.modelResponse(ctx, input) return imageBundle, false, modelErr } @@ -94,7 +98,9 @@ func (runtime *Runtime) respond( }, true, nil } - bundle, err = runtime.modelResponse(ctx, sessionID, lineage, cwd, prompt, false, onEvent, onRetry) + input.contextHasImages = false + + bundle, err = runtime.modelResponse(ctx, input) if err != nil { return nil, false, err } @@ -106,13 +112,7 @@ func (runtime *Runtime) respond( func (runtime *Runtime) modelResponse( ctx context.Context, - sessionID string, - lineage *promptLineage, - cwd string, - prompt string, - contextHasImages bool, - onEvent func(StreamEvent), - onRetry RetryEventHandler, + input *responseInput, ) (*responseBundle, error) { if runtime.models == nil { return nil, oops.In("assistant").Code("models_unavailable").Errorf("model registry is not configured") @@ -123,7 +123,7 @@ func (runtime *Runtime) modelResponse( return nil, err } - if contextHasImages { + if input.contextHasImages { // Keep this early check: historical images are known before auth and context // construction, so an incompatible model should fail without doing either. imageErr := validateSelectedModelHasImageInput(&selectedModel, "conversation_history") @@ -141,13 +141,13 @@ func (runtime *Runtime) modelResponse( } preparation := &completionRequestPreparationInput{ - sessionID: sessionID, - cwd: cwd, - prompt: prompt, - lineage: lineage, + sessionID: input.sessionID, + cwd: input.cwd, + prompt: input.prompt, + lineage: input.lineage, selectedModel: &selectedModel, auth: &auth, - onEvent: onEvent, + onEvent: input.onEvent, } build, compactionEntry, err := runtime.prepareCompletionRequestWithAutoCompaction(ctx, preparation) @@ -168,7 +168,7 @@ func (runtime *Runtime) modelResponse( preparation: preparation, build: build, compactionEntry: compactionEntry, - onRetry: onRetry, + onRetry: input.onRetry, }, ) if err != nil { @@ -176,9 +176,9 @@ func (runtime *Runtime) modelResponse( } usage := contextwindow.MergeUsage(build.Context.Usage, result.Usage) - runtime.emitUsage(ctx, onEvent, usage) + runtime.emitUsage(ctx, input.onEvent, usage) - lineage.adopt(compactionEntry) + input.lineage.adopt(compactionEntry) return &responseBundle{ Text: result.Text, diff --git a/internal/assistant/runtime_persist.go b/internal/assistant/runtime_persist.go index c45e2dc4..35a43b63 100644 --- a/internal/assistant/runtime_persist.go +++ b/internal/assistant/runtime_persist.go @@ -106,16 +106,16 @@ func (runtime *Runtime) respondWithPartialProgress( ) (*responseBundle, bool, error) { progress := newPartialPromptProgress(request.OnEvent) - bundle, cached, err := runtime.respond( - ctx, - sessionID, - lineage, - request.CWD, - request.Text, - len(request.Images) > 0, - progress.handle, - progress.retryHandler(request.OnRetry), - ) + bundle, cached, err := runtime.respond(ctx, &responseInput{ + lineage: lineage, + onEvent: progress.handle, + onRetry: progress.retryHandler(request.OnRetry), + sessionID: sessionID, + cwd: request.CWD, + prompt: request.Text, + hasPromptImages: len(request.Images) > 0, + contextHasImages: false, + }) if err != nil { persistErr := runtime.appendPartialPromptFailure( ctx, diff --git a/internal/database/migrations/00013_add_session_message_parts.sql b/internal/database/migrations/00013_add_session_message_parts.sql index 7a4c2064..9cd3f97f 100644 --- a/internal/database/migrations/00013_add_session_message_parts.sql +++ b/internal/database/migrations/00013_add_session_message_parts.sql @@ -1,12 +1,12 @@ -- +goose Up -CREATE UNIQUE INDEX IF NOT EXISTS idx_session_entries_id_session +CREATE UNIQUE INDEX idx_session_entries_id_session ON session_entries(id, session_id); -CREATE TABLE IF NOT EXISTS session_message_parts ( - id TEXT PRIMARY KEY, +CREATE TABLE session_message_parts ( + id TEXT NOT NULL PRIMARY KEY, session_id TEXT NOT NULL, entry_id TEXT NOT NULL, - sequence INTEGER NOT NULL CHECK(sequence >= 0), + sequence INTEGER NOT NULL CHECK(typeof(sequence) = 'integer' AND sequence >= 0), type TEXT NOT NULL CHECK(type IN ('text', 'image')), text TEXT NOT NULL DEFAULT '', mime_type TEXT NOT NULL DEFAULT '', @@ -19,7 +19,7 @@ CREATE TABLE IF NOT EXISTS session_message_parts ( UNIQUE (entry_id, sequence) ); -CREATE INDEX IF NOT EXISTS idx_session_message_parts_session_entry_sequence +CREATE INDEX idx_session_message_parts_session_entry_sequence ON session_message_parts(session_id, entry_id, sequence); -- +goose Down diff --git a/internal/database/migrations/00014_index_image_message_parts.sql b/internal/database/migrations/00014_index_image_message_parts.sql index 73a8ccd3..e07ccf8d 100644 --- a/internal/database/migrations/00014_index_image_message_parts.sql +++ b/internal/database/migrations/00014_index_image_message_parts.sql @@ -1,5 +1,5 @@ -- +goose Up -CREATE INDEX IF NOT EXISTS idx_session_message_parts_session_entry_image +CREATE INDEX idx_session_message_parts_session_entry_image ON session_message_parts(session_id, entry_id) WHERE type = 'image'; diff --git a/internal/database/migrations_test.go b/internal/database/migrations_test.go index 9be76d95..a89940ed 100644 --- a/internal/database/migrations_test.go +++ b/internal/database/migrations_test.go @@ -5,6 +5,7 @@ import ( "database/sql" "testing" "testing/fstest" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -12,7 +13,10 @@ import ( "github.com/omarluq/librecode/internal/database" ) -const deployedWorkflowMigrationV8 = `-- +goose Up +const ( + schemaIndexType = "index" + + deployedWorkflowMigrationV8 = `-- +goose Up CREATE TABLE workflow_runs ( task_id TEXT PRIMARY KEY, source TEXT NOT NULL, @@ -38,6 +42,7 @@ CREATE INDEX idx_workflow_agent_tasks_replay DROP TABLE IF EXISTS workflow_agent_tasks; DROP TABLE IF EXISTS workflow_runs; ` +) func TestMessagePartsMigrationUpDownAndOldSchemaUpgrade(t *testing.T) { t.Parallel() @@ -65,6 +70,24 @@ func TestMessagePartsMigrationUpDownAndOldSchemaUpgrade(t *testing.T) { `SELECT COUNT(*) FROM sqlite_master WHERE type = 'index' AND name = ?`, indexName) } + repository := database.NewSessionRepository(connection) + session, err := repository.CreateSession(ctx, "/work", "parts constraints", "") + require.NoError(t, err) + entry, err := repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: testMessageText, + Provider: "", Model: "", Parts: nil, + }) + require.NoError(t, err) + + _, err = connection.ExecContext(ctx, ` +INSERT INTO session_message_parts (id, session_id, entry_id, sequence, type, text) +VALUES (NULL, ?, ?, 1, 'text', 'invalid')`, session.ID, entry.ID) + require.ErrorContains(t, err, "NOT NULL constraint failed") + _, err = connection.ExecContext(ctx, ` +INSERT INTO session_message_parts (id, session_id, entry_id, sequence, type, text) +VALUES ('fractional-sequence', ?, ?, 1.5, 'text', 'invalid')`, session.ID, entry.ID) + require.ErrorContains(t, err, "CHECK constraint failed") + _, err = provider.Down(ctx) require.NoError(t, err) assertSchemaObjectExists(ctx, t, connection, @@ -80,6 +103,70 @@ WHERE type = 'index' AND name = 'idx_session_message_parts_session_entry_image'` `SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'session_message_parts'`) } +func TestMessagePartsMigrationRejectsPreexistingSchemaObjects(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + setup string + objectType string + objectName string + }{ + { + name: "index", + setup: `CREATE INDEX idx_session_entries_id_session ON session_entries(session_id)`, + objectType: schemaIndexType, + objectName: "idx_session_entries_id_session", + }, + { + name: "table", + setup: `CREATE TABLE session_message_parts (sentinel TEXT NOT NULL)`, + objectType: "table", + objectName: "session_message_parts", + }, + { + name: "image index", + setup: ` +CREATE UNIQUE INDEX idx_session_entries_id_session ON session_entries(id, session_id); +CREATE TABLE session_message_parts ( + id TEXT NOT NULL PRIMARY KEY, + session_id TEXT NOT NULL, + entry_id TEXT NOT NULL, + sequence INTEGER NOT NULL, + type TEXT NOT NULL +); +CREATE INDEX idx_session_message_parts_session_entry_sequence + ON session_message_parts(session_id, entry_id, sequence); +INSERT INTO goose_db_version (version_id, is_applied) VALUES (13, 1); +CREATE INDEX idx_session_message_parts_session_entry_image + ON session_message_parts(entry_id)`, + objectType: schemaIndexType, + objectName: "idx_session_message_parts_session_entry_image", + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + connection := newMigratedThroughVersion(t, 12) + ctx := context.Background() + _, err := connection.ExecContext(ctx, test.setup) + require.NoError(t, err) + + migrationRoot, migrationErr := database.MigrationFS() + require.NoError(t, migrationErr) + provider, migrationErr := database.NewMigrationProvider(connection, migrationRoot) + require.NoError(t, migrationErr) + _, migrationErr = provider.Up(ctx) + require.Error(t, migrationErr) + + assertSchemaObjectExists(ctx, t, connection, + `SELECT COUNT(*) FROM sqlite_master WHERE type = ? AND name = ?`, + test.objectType, test.objectName) + }) + } +} + func assertSchemaObjectCount( ctx context.Context, t *testing.T, diff --git a/internal/database/session_entry_repository.go b/internal/database/session_entry_repository.go index 1bba76b1..ae3ed434 100644 --- a/internal/database/session_entry_repository.go +++ b/internal/database/session_entry_repository.go @@ -92,26 +92,7 @@ WHERE session_id = ? ORDER BY created_at DESC LIMIT 1`, entrySelectColumns) - var row entryRow - if err := repository.sql.QueryOne(ctx, &row, query, sessionID); err != nil { - if errors.Is(err, ksql.ErrRecordNotFound) { - return nil, false, nil - } - - return nil, false, oops.In("database").Code("leaf_entry").Wrapf(err, "load leaf entry") - } - - entry, err := entryFromRow(&row) - if err != nil { - return nil, false, oops.In("database").Code("scan_entry").Wrapf(err, "scan leaf entry") - } - - entries := []EntryEntity{*entry} - if err := repository.hydrateEntryMessages(ctx, sessionID, entries); err != nil { - return nil, false, err - } - - return &entries[0], true, nil + return repository.queryEntry(ctx, sessionID, query, "leaf_entry", "load leaf entry", sessionID) } // Entries returns all entries for a session in append order. @@ -146,13 +127,24 @@ SELECT %s FROM session_entries WHERE session_id = ? AND id = ?`, entrySelectColumns) + return repository.queryEntry(ctx, sessionID, query, "get_entry", "load entry", sessionID, entryID) +} + +func (repository *SessionRepository) queryEntry( + ctx context.Context, + sessionID string, + query string, + queryCode string, + queryMessage string, + arguments ...any, +) (*EntryEntity, bool, error) { var row entryRow - if err := repository.sql.QueryOne(ctx, &row, query, sessionID, entryID); err != nil { + if err := repository.sql.QueryOne(ctx, &row, query, arguments...); err != nil { if errors.Is(err, ksql.ErrRecordNotFound) { return nil, false, nil } - return nil, false, oops.In("database").Code("get_entry").Wrapf(err, "load entry") + return nil, false, oops.In("database").Code(queryCode).Wrapf(err, "%s", queryMessage) } entry, err := entryFromRow(&row) @@ -160,12 +152,11 @@ WHERE session_id = ? AND id = ?`, entrySelectColumns) return nil, false, oops.In("database").Code("scan_entry").Wrapf(err, "scan entry") } - entries := []EntryEntity{*entry} - if err := repository.hydrateEntryMessages(ctx, sessionID, entries); err != nil { + if err := repository.hydrateEntryMessage(ctx, sessionID, entry); err != nil { return nil, false, err } - return &entries[0], true, nil + return entry, true, nil } // DeleteEntryBranch removes an entry and all descendants from one session. diff --git a/internal/database/session_message_parts_test.go b/internal/database/session_message_parts_test.go index f804ac2a..c41cbd43 100644 --- a/internal/database/session_message_parts_test.go +++ b/internal/database/session_message_parts_test.go @@ -12,7 +12,11 @@ import ( "github.com/omarluq/librecode/internal/database" ) -const testImageMIME = "image/png" +const ( + testImageMIME = "image/png" + testMessageText = "text" + testMultipartText = "compare these" +) func TestSessionRepository_RoundTripsOrderedMultipartMessages(t *testing.T) { t.Parallel() @@ -24,7 +28,10 @@ func TestSessionRepository_RoundTripsOrderedMultipartMessages(t *testing.T) { originalData := []byte{1, 2, 3} parts := []database.MessagePartEntity{ - {Data: nil, Text: "compare these", MIMEType: "", Name: "", Type: database.MessagePartText, Width: 0, Height: 0}, + { + Data: nil, Text: testMultipartText, MIMEType: "", Name: "", + Type: database.MessagePartText, Width: 0, Height: 0, + }, { Data: originalData, Text: "", MIMEType: testImageMIME, Name: "first.png", Type: database.MessagePartImage, Width: 10, Height: 20, @@ -163,48 +170,111 @@ func TestSessionRepository_PersistsImageOnlyAndSupportsLegacyText(t *testing.T) assert.Equal(t, []byte{7}, contextEntity.Messages[1].Parts[0].Data) } -func TestSessionRepository_RejectsMultipartResourceLimitBypasses(t *testing.T) { +func TestSessionRepository_RejectsInvalidMultipartMessages(t *testing.T) { t.Parallel() - repository := newTestSessionRepository(t) - ctx := context.Background() - session, err := repository.CreateSession(ctx, "/work", "limits", "") - require.NoError(t, err) + for _, test := range invalidMultipartMessageCases() { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + repository := newTestSessionRepository(t) + session, err := repository.CreateSession(t.Context(), "/work", test.name, "") + require.NoError(t, err) + _, err = repository.AppendMessage(t.Context(), session.ID, nil, &database.MessageEntity{ + Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: test.content, Provider: "", Model: "", + Parts: test.parts, + }) + require.ErrorContains(t, err, test.wantErr) + + entries, err := repository.Entries(t.Context(), session.ID) + require.NoError(t, err) + assert.Empty(t, entries) + }) + } +} + +type invalidMultipartMessageCase struct { + name string + content string + wantErr string + parts []database.MessagePartEntity +} +func invalidMultipartMessageCases() []invalidMultipartMessageCase { images := make([]database.MessagePartEntity, 5) for index := range images { - images[index] = database.MessagePartEntity{ - Text: "", MIMEType: testImageMIME, Name: "", Type: database.MessagePartImage, - Data: []byte{1}, Width: 1, Height: 1, - } + images[index] = testImagePart([]byte{1}, "", 1, 1) } - _, err = repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ - Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "", Provider: "", Model: "", Parts: images, - }) - require.ErrorContains(t, err, "maximum is 4") - - _, err = repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ - Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "", Provider: "", Model: "", - Parts: []database.MessagePartEntity{{ - Text: "", MIMEType: "image/*", Name: "", Type: database.MessagePartImage, - Data: []byte{1}, Width: 1, Height: 1, - }}, - }) - require.ErrorContains(t, err, "normalized image MIME type") + return []invalidMultipartMessageCase{ + {name: "image count", content: "", wantErr: "maximum is 4", parts: images}, + { + name: "text projection mismatch", content: "different", + wantErr: "content must match the text projection", + parts: []database.MessagePartEntity{testTextPart(testMessageText, nil)}, + }, + { + name: "unsupported type", content: "", wantErr: "unsupported message part type", + parts: []database.MessagePartEntity{{ + Data: nil, Text: "", MIMEType: "", Name: "", + Type: database.MessagePartType("audio"), Width: 0, Height: 0, + }}, + }, + { + name: "blank text", content: "", wantErr: "text part must have text", + parts: []database.MessagePartEntity{testTextPart(" \t", nil)}, + }, + { + name: "text with binary data", content: testMessageText, + wantErr: "text part must not have binary data", + parts: []database.MessagePartEntity{testTextPart(testMessageText, []byte{1})}, + }, + { + name: "image without binary data", content: "", wantErr: "image part must have binary data", + parts: []database.MessagePartEntity{testImagePart(nil, "", 1, 1)}, + }, + { + name: "image with text", content: "", wantErr: "image part must not have text", + parts: []database.MessagePartEntity{testImagePart([]byte{1}, testMessageText, 1, 1)}, + }, + { + name: "MIME type", content: "", wantErr: "normalized image MIME type", + parts: []database.MessagePartEntity{{ + Data: []byte{1}, Text: "", MIMEType: "image/*", Name: "", + Type: database.MessagePartImage, Width: 1, Height: 1, + }}, + }, + { + name: "image byte size", content: "", wantErr: "5 MiB limit", + parts: []database.MessagePartEntity{testImagePart(make([]byte, 5*1024*1024+1), "", 1, 1)}, + }, + { + name: "zero width", content: "", wantErr: "dimensions must be positive", + parts: []database.MessagePartEntity{testImagePart([]byte{1}, "", 0, 1)}, + }, + { + name: "zero height", content: "", wantErr: "dimensions must be positive", + parts: []database.MessagePartEntity{testImagePart([]byte{1}, "", 1, 0)}, + }, + { + name: "pixel count", content: "", wantErr: "40 megapixels", + parts: []database.MessagePartEntity{testImagePart([]byte{1}, "", 40_000_001, 1)}, + }, + } +} - _, err = repository.AppendMessage(ctx, session.ID, nil, &database.MessageEntity{ - Timestamp: time.Now().UTC(), Role: database.RoleUser, Content: "", Provider: "", Model: "", - Parts: []database.MessagePartEntity{{ - Text: "", MIMEType: testImageMIME, Name: "", Type: database.MessagePartImage, - Data: []byte{1}, Width: 40_000_001, Height: 1, - }}, - }) - require.ErrorContains(t, err, "40 megapixels") +func testTextPart(text string, data []byte) database.MessagePartEntity { + return database.MessagePartEntity{ + Data: data, Text: text, MIMEType: "", Name: "", + Type: database.MessagePartText, Width: 0, Height: 0, + } +} - entries, listErr := repository.Entries(ctx, session.ID) - require.NoError(t, listErr) - assert.Empty(t, entries) +func testImagePart(data []byte, text string, width, height int) database.MessagePartEntity { + return database.MessagePartEntity{ + Data: data, Text: text, MIMEType: testImageMIME, Name: "", + Type: database.MessagePartImage, Width: width, Height: height, + } } func TestSessionRepository_PartInsertFailureRollsBackEntryAndMessage(t *testing.T) { @@ -336,8 +406,10 @@ func assertEntryReadersHydrateMultipart( func assertMultipartParts(t *testing.T, parts []database.MessagePartEntity) { t.Helper() require.Len(t, parts, 3) - assert.Equal(t, database.MessagePartText, parts[0].Type) - assert.Equal(t, "compare these", parts[0].Text) + assert.Equal(t, database.MessagePartEntity{ + Data: nil, Text: testMultipartText, MIMEType: "", Name: "", + Type: database.MessagePartText, Width: 0, Height: 0, + }, parts[0]) assert.Equal(t, database.MessagePartImage, parts[1].Type) assert.Equal(t, []byte{1, 2, 3}, parts[1].Data) assert.Equal(t, "first.png", parts[1].Name) diff --git a/internal/provider/image_content.go b/internal/provider/image_content.go index 153828fb..dc9d3cf4 100644 --- a/internal/provider/image_content.go +++ b/internal/provider/image_content.go @@ -98,34 +98,12 @@ func openAIResponseUserContent(message llm.Message) []map[string]any { } func openAIChatUserContent(message llm.Message) any { - hasImage := false - for index := range message.Content { - hasImage = hasImage || message.Content[index].Type == llm.PartImage - } - - if !hasImage { - return messageText(message) - } - - blocks := make([]map[string]any, 0, len(message.Content)) - for index := range message.Content { - part := &message.Content[index] - switch part.Type { - case llm.PartText: - if part.Text != "" { - blocks = append(blocks, map[string]any{jsonTypeKey: jsonTextKey, jsonTextKey: part.Text}) - } - case llm.PartImage: - blocks = append(blocks, map[string]any{ - jsonTypeKey: jsonImageURLKey, - jsonImageURLKey: map[string]any{jsonImageURLValueKey: openAIDataURL(part)}, - }) - case llm.PartReasoning, llm.PartFile, llm.PartSource, llm.PartToolCall, llm.PartToolResult: - continue + return structuredUserContent(message, func(part *llm.Part) map[string]any { + return map[string]any{ + jsonTypeKey: jsonImageURLKey, + jsonImageURLKey: map[string]any{jsonImageURLValueKey: openAIDataURL(part)}, } - } - - return blocks + }) } func emptyMessageContent(content any) bool { @@ -140,6 +118,14 @@ func emptyMessageContent(content any) bool { } func anthropicUserContent(message llm.Message) any { + return structuredUserContent(message, func(part *llm.Part) map[string]any { + return map[string]any{jsonTypeKey: anthropicImageType, anthropicSourceKey: map[string]any{ + jsonTypeKey: anthropicBase64Type, anthropicMediaTypeKey: part.MIMEType, encodedImageDataKey: part.Data, + }} + }) +} + +func structuredUserContent(message llm.Message, imageBlock func(*llm.Part) map[string]any) any { hasImage := false for index := range message.Content { hasImage = hasImage || message.Content[index].Type == llm.PartImage @@ -158,9 +144,7 @@ func anthropicUserContent(message llm.Message) any { blocks = append(blocks, map[string]any{jsonTypeKey: jsonTextKey, jsonTextKey: part.Text}) } case llm.PartImage: - blocks = append(blocks, map[string]any{jsonTypeKey: anthropicImageType, anthropicSourceKey: map[string]any{ - jsonTypeKey: anthropicBase64Type, anthropicMediaTypeKey: part.MIMEType, encodedImageDataKey: part.Data, - }}) + blocks = append(blocks, imageBlock(part)) case llm.PartReasoning, llm.PartFile, llm.PartSource, llm.PartToolCall, llm.PartToolResult: continue } diff --git a/internal/terminal/compact_commands_internal_test.go b/internal/terminal/compact_commands_internal_test.go index e4da5ecb..5559e4b7 100644 --- a/internal/terminal/compact_commands_internal_test.go +++ b/internal/terminal/compact_commands_internal_test.go @@ -202,6 +202,8 @@ func TestHandleCompactDoneStartsQueuedPrompt(t *testing.T) { message := request.Messages[len(request.Messages)-1] assert.Equal(t, "queued after compact", message.Content) require.Len(t, message.Parts, 2) + assert.Equal(t, database.MessagePartText, message.Parts[0].Type) + assert.Equal(t, "queued after compact", message.Parts[0].Text) assert.Equal(t, database.MessagePartImage, message.Parts[1].Type) assert.Equal(t, image.Name, message.Parts[1].Name) assert.Equal(t, image.MIMEType, message.Parts[1].MIMEType) From 69f5c7a7952b1b7f953ebaa993ca7f717f9a97d7 Mon Sep 17 00:00:00 2001 From: Omar Alani Date: Tue, 4 Aug 2026 10:23:42 -0500 Subject: [PATCH 4/6] fix(deps): update image library for security fixes --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/go.mod b/go.mod index d9444e06..d6e9ef9b 100644 --- a/go.mod +++ b/go.mod @@ -48,7 +48,7 @@ require ( github.com/yuin/goldmark v1.8.5 github.com/yuin/gopher-lua v1.1.2 golang.design/x/clipboard v0.8.0 - golang.org/x/image v0.41.0 + golang.org/x/image v0.44.0 golang.org/x/net v0.57.0 golang.org/x/text v0.40.0 gopkg.in/yaml.v3 v3.0.1 diff --git a/go.sum b/go.sum index 2464726e..d2338706 100644 --- a/go.sum +++ b/go.sum @@ -267,8 +267,8 @@ golang.org/x/exp v0.0.0-20260718201538-764159d718ef h1:LkZ48HFgy/TvhTI0bcWkjgFkg golang.org/x/exp v0.0.0-20260718201538-764159d718ef/go.mod h1:EdfpwwqSu+0Li0mzskwHU6FWDV3t9Q+RZDo3QMUtL3Q= golang.org/x/exp/shiny v0.0.0-20250606033433-dcc06ee1d476 h1:Wdx0vgH5Wgsw+lF//LJKmWOJBLWX6nprsMqnf99rYDE= golang.org/x/exp/shiny v0.0.0-20250606033433-dcc06ee1d476/go.mod h1:ygj7T6vSGhhm/9yTpOQQNvuAUFziTH7RUiH74EoE2C8= -golang.org/x/image v0.41.0 h1:8wS72eGJMJaBxK6okTzd4WaXumUlTVlb753MlsSvTCo= -golang.org/x/image v0.41.0/go.mod h1:uIc348UZMSvS5Z65CVZ7iDPaNobNFEPeJ4kbqTOszmA= +golang.org/x/image v0.44.0 h1:+tDekMZED9+LrtB3G5xzRggpVh9CARjZqROla3R3R+I= +golang.org/x/image v0.44.0/go.mod h1:V8K3KE9KKKE+pLpQDOeN18w9oacNSvy1tDOirTu4xtY= golang.org/x/mobile v0.0.0-20250606033058-a2a15c67f36f h1:/n+PL2HlfqeSiDCuhdBbRNlGS/g2fM4OHufalHaTVG8= golang.org/x/mobile v0.0.0-20250606033058-a2a15c67f36f/go.mod h1:ESkJ836Z6LpG6mTVAhA48LpfW/8fNR0ifStlH2axyfg= golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= From 49cf7d4b66a9e2c5c15e9edd360a09c224cf4c8e Mon Sep 17 00:00:00 2001 From: Omar Alani Date: Tue, 4 Aug 2026 10:53:27 -0500 Subject: [PATCH 5/6] fix(terminal): preserve composer borders around attachments --- internal/terminal/attachments_internal_test.go | 3 +++ internal/terminal/render_composer.go | 14 ++++++++++---- 2 files changed, 13 insertions(+), 4 deletions(-) diff --git a/internal/terminal/attachments_internal_test.go b/internal/terminal/attachments_internal_test.go index a5f9b9e5..47eb4603 100644 --- a/internal/terminal/attachments_internal_test.go +++ b/internal/terminal/attachments_internal_test.go @@ -260,8 +260,11 @@ func TestAttachmentRenderingContainsMetadataNotData(t *testing.T) { attachment := imageAttachment{Name: "paste-1.png", MIMEType: "image/png", Data: secret, Width: 10, Height: 20} app.composerImages = []imageAttachment{attachment} composer := app.attachmentChipLines(80) + require.Len(t, composer, 1) assert.Contains(t, composer[0].Text, "paste-1.png") assert.NotContains(t, composer[0].Text, string(secret)) + assert.True(t, strings.HasPrefix(composer[0].Text, "│ ")) + assert.True(t, strings.HasSuffix(composer[0].Text, " │")) lines := app.renderUserMessage(80, "", *summarizeAttachments([]imageAttachment{attachment})) diff --git a/internal/terminal/render_composer.go b/internal/terminal/render_composer.go index 714d9701..a90e88c8 100644 --- a/internal/terminal/render_composer.go +++ b/internal/terminal/render_composer.go @@ -73,13 +73,19 @@ func (app *App) attachmentChipLines(width int) []tui.Line { Name: item.Name, MIMEType: item.MIMEType, Width: item.Width, Height: item.Height, Size: len(item.Data), } - text := " " + attachmentSummaryText(summary) - lines = append(lines, tui.NewLine(app.theme.style(colorDim), tui.Truncate(text, width))) + lines = append(lines, app.attachmentChipLine(width, attachmentSummaryText(summary))) } return lines } +func (app *App) attachmentChipLine(width int, text string) tui.Line { + innerWidth := max(1, width-terminalMarkerMargin*2) + body := tui.PadRight(tui.Truncate(text, innerWidth), innerWidth) + + return tui.NewLine(app.theme.style(colorDim), "│ "+body+" │") +} + func (app *App) visibleAttachmentChipLines(width, limit int) []tui.Line { chips := app.attachmentChipLines(width) if limit <= 0 || len(chips) == 0 { @@ -92,9 +98,9 @@ func (app *App) visibleAttachmentChipLines(width, limit int) []tui.Line { visible := max(0, limit-1) lines := chips[:visible] - overflow := " … " + tui.Int(len(chips)-visible) + " more attachments" + overflow := "… " + tui.Int(len(chips)-visible) + " more attachments" - return append(lines, tui.NewLine(app.theme.style(colorDim), tui.Truncate(overflow, width))) + return append(lines, app.attachmentChipLine(width, overflow)) } func formatByteSize(size int) string { From 0cba29d57ee051c1062293685cb2ebd282723f27 Mon Sep 17 00:00:00 2001 From: Omar Alani Date: Tue, 4 Aug 2026 10:53:33 -0500 Subject: [PATCH 6/6] test(database): verify persisted image metadata --- internal/database/session_message_parts_test.go | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/internal/database/session_message_parts_test.go b/internal/database/session_message_parts_test.go index c41cbd43..00cb9207 100644 --- a/internal/database/session_message_parts_test.go +++ b/internal/database/session_message_parts_test.go @@ -410,9 +410,12 @@ func assertMultipartParts(t *testing.T, parts []database.MessagePartEntity) { Data: nil, Text: testMultipartText, MIMEType: "", Name: "", Type: database.MessagePartText, Width: 0, Height: 0, }, parts[0]) - assert.Equal(t, database.MessagePartImage, parts[1].Type) - assert.Equal(t, []byte{1, 2, 3}, parts[1].Data) - assert.Equal(t, "first.png", parts[1].Name) - assert.Equal(t, database.MessagePartImage, parts[2].Type) - assert.Equal(t, []byte{4, 5}, parts[2].Data) + assert.Equal(t, database.MessagePartEntity{ + Data: []byte{1, 2, 3}, Text: "", MIMEType: testImageMIME, Name: "first.png", + Type: database.MessagePartImage, Width: 10, Height: 20, + }, parts[1]) + assert.Equal(t, database.MessagePartEntity{ + Data: []byte{4, 5}, Text: "", MIMEType: "image/jpeg", Name: "second.jpg", + Type: database.MessagePartImage, Width: 30, Height: 40, + }, parts[2]) }