From 072c0158ffcbc2831360430188b22b2db76e3401 Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Sat, 6 Jun 2026 11:50:24 -0700 Subject: [PATCH] feat: support dashboard chat images (#1435) * feat: support dashboard chat images Co-authored-by: open-swe[bot] * fix: keep pending image prompts visible Co-authored-by: open-swe[bot] --------- Co-authored-by: open-swe[bot] --- agent/dashboard/message_adapter.py | 49 ++++- agent/dashboard/thread_api.py | 78 +++++++- agent/middleware/check_message_queue.py | 7 +- tests/test_dashboard_message_adapter.py | 32 +++ tests/test_dashboard_web_handoff.py | 87 ++++++++ ui/src/components/agents/AgentThreadView.tsx | 54 +++-- ui/src/components/agents/AgentsHome.tsx | 3 +- .../agents/ported/CloudPromptBar.tsx | 187 ++++++++++++++++-- ui/src/lib/agents/api.ts | 4 +- ui/src/lib/agents/pendingPrompts.ts | 42 +++- ui/src/lib/agents/queries.ts | 12 +- 11 files changed, 503 insertions(+), 52 deletions(-) diff --git a/agent/dashboard/message_adapter.py b/agent/dashboard/message_adapter.py index d6314eca..6a5b61a6 100644 --- a/agent/dashboard/message_adapter.py +++ b/agent/dashboard/message_adapter.py @@ -3,6 +3,7 @@ from __future__ import annotations import json +import re import uuid from datetime import UTC, datetime from typing import Any @@ -14,6 +15,7 @@ _EDIT_TOOLS = frozenset({"write_file", "edit_file", "str_replace", "write", "edi _EXECUTE_TOOLS = frozenset({"execute", "bash", "shell", "run_terminal_cmd"}) _SEARCH_TOOLS = frozenset({"glob", "grep", "web_search", "fetch_url", "search"}) _INTERNAL_TOOLS = frozenset({"confirming_completion", "no_op"}) +_DATA_IMAGE_RE = re.compile(r"^data:(image/[^;]+);base64,(.+)$", re.DOTALL) def _now_iso() -> str: @@ -78,6 +80,43 @@ def _parse_tool_args(raw: Any) -> dict[str, Any]: return {} +def _image_chunks(content: Any) -> list[dict[str, Any]]: + if not isinstance(content, list): + return [] + + chunks: list[dict[str, Any]] = [] + for item in content: + if not isinstance(item, dict): + continue + item_type = item.get("type") + base64_data: str | None = None + mime_type: str | None = None + if item_type == "image": + data = item.get("data") or item.get("base64") + mime = item.get("mime_type") or item.get("mimeType") + if isinstance(data, str) and isinstance(mime, str): + base64_data = data + mime_type = mime + elif item_type == "image_url": + image_url = item.get("image_url") + url = image_url.get("url") if isinstance(image_url, dict) else None + if isinstance(url, str): + match = _DATA_IMAGE_RE.match(url) + if match: + mime_type, base64_data = match.groups() + if base64_data and mime_type: + chunk: dict[str, Any] = { + "kind": "image", + "base64": base64_data, + "mimeType": mime_type, + } + file_name = item.get("fileName") or item.get("file_name") + if isinstance(file_name, str) and file_name: + chunk["fileName"] = file_name + chunks.append(chunk) + return chunks + + def _maybe_diff_from_args(name: str, args: dict[str, Any]) -> dict[str, Any] | None: path = args.get("path") or args.get("file_path") or args.get("target_file") if not isinstance(path, str) or not path.strip(): @@ -150,15 +189,19 @@ def state_messages_to_ui(messages: list[Any]) -> list[dict[str, Any]]: agent_turn["chunks"] = _merge_text_chunks(agent_turn["chunks"]) ui_messages.append(agent_turn) agent_turn = None - text = extract_text_content(raw.get("content", "")) - if not text: + content = raw.get("content", "") + chunks = _image_chunks(content) + text = extract_text_content(content) + if text: + chunks.append({"kind": "text", "text": text}) + if not chunks: continue ui_messages.append( { "id": msg_id, "author": "user", "timestamp": timestamp, - "chunks": [{"kind": "text", "text": text}], + "chunks": chunks, } ) continue diff --git a/agent/dashboard/thread_api.py b/agent/dashboard/thread_api.py index 43be0ca6..da906ee8 100644 --- a/agent/dashboard/thread_api.py +++ b/agent/dashboard/thread_api.py @@ -2,6 +2,8 @@ from __future__ import annotations +import base64 +import binascii import json import logging import os @@ -11,8 +13,9 @@ from datetime import UTC, datetime from typing import Any from fastapi import HTTPException +from langchain_core.messages.content import create_image_block from langgraph_sdk.errors import InternalServerError -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field from ..utils.thread_ops import is_thread_active, langgraph_client, queue_message_for_thread from .message_adapter import state_messages_to_ui @@ -25,6 +28,9 @@ logger = logging.getLogger(__name__) _ASSISTANT_ID = "agent" _DASHBOARD_SOURCE = "dashboard" _DASHBOARD_STREAM_MODES: tuple[str, ...] = ("values", "updates", "messages-tuple") +_SUPPORTED_IMAGE_MIME_TYPES = frozenset({"image/png", "image/jpeg", "image/gif", "image/webp"}) +_MAX_DASHBOARD_IMAGES = 5 +_MAX_DASHBOARD_IMAGE_BYTES = 10 * 1024 * 1024 # Sources whose threads should surface in the Agents UI (besides "dashboard"). _SURFACED_SOURCES: tuple[str, ...] = ("dashboard", "github", "slack", "linear", "schedule") @@ -45,8 +51,18 @@ async def _resolve_run_email(login: str, profile: dict[str, Any]) -> str | None: return mapped or profile.get("email") +class DashboardImageBody(BaseModel): + model_config = ConfigDict(populate_by_name=True) + + kind: str | None = None + base64: str = Field(min_length=1) + mime_type: str = Field(alias="mimeType", min_length=1) + file_name: str | None = Field(default=None, alias="fileName") + + class ThreadCreateBody(BaseModel): - prompt: str = Field(min_length=1, max_length=20_000) + prompt: str = Field(default="", max_length=20_000) + images: list[DashboardImageBody] = Field(default_factory=list) repo: str | None = None repo_explicitly_none: bool = False model_id: str | None = None @@ -54,7 +70,8 @@ class ThreadCreateBody(BaseModel): class ThreadMessageBody(BaseModel): - content: str = Field(min_length=1, max_length=20_000) + content: str = Field(default="", max_length=20_000) + images: list[DashboardImageBody] = Field(default_factory=list) model_id: str | None = None effort: str | None = None @@ -85,6 +102,41 @@ def _parse_repo(full_name: str | None) -> dict[str, str] | None: return {"owner": owner, "name": name} +def _decode_dashboard_image(image: DashboardImageBody) -> bytes: + if image.mime_type not in _SUPPORTED_IMAGE_MIME_TYPES: + raise HTTPException(422, f"unsupported image type: {image.mime_type}") + try: + data = base64.b64decode(image.base64, validate=True) + except binascii.Error as exc: + raise HTTPException(422, "invalid image data") from exc + if len(data) > _MAX_DASHBOARD_IMAGE_BYTES: + raise HTTPException(422, "image exceeds 10MB limit") + return data + + +def _image_blocks(images: list[DashboardImageBody]) -> list[dict[str, Any]]: + if len(images) > _MAX_DASHBOARD_IMAGES: + raise HTTPException(422, f"at most {_MAX_DASHBOARD_IMAGES} images are supported") + return [ + create_image_block( + base64=base64.b64encode(_decode_dashboard_image(image)).decode("ascii"), + mime_type=image.mime_type, + ) + for image in images + ] + + +def _user_message_content( + prompt: str, images: list[DashboardImageBody] +) -> str | list[dict[str, Any]]: + text = prompt.strip() + if not text and not images: + raise HTTPException(422, "prompt or image required") + if not images: + return text + return [*_image_blocks(images), *([{"type": "text", "text": text}] if text else [])] + + async def _ensure_dashboard_github_token(login: str) -> None: token = await get_valid_access_token(login) if not token: @@ -283,12 +335,15 @@ async def _start_agent_run( repo_config: dict[str, str], repo_explicitly_none: bool = False, prompt: str, + images: list[DashboardImageBody] | None = None, title: str | None = None, model_id: str | None = None, effort: str | None = None, ) -> dict[str, Any]: profile = await get_profile(login) or {} now_ms = _now_ms() + prompt = prompt.strip() + content = _user_message_content(prompt, images or []) chosen_model, chosen_effort = _normalize_model_choice(model_id, effort) metadata_model = chosen_model or profile.get("default_model") or "Default" metadata_effort = chosen_effort or profile.get("reasoning_effort") @@ -332,7 +387,7 @@ async def _start_agent_run( run = await client.runs.create( thread_id, _ASSISTANT_ID, - input={"messages": [{"role": "user", "content": prompt}]}, + input={"messages": [{"role": "user", "content": content}]}, config={"configurable": configurable, "metadata": _agent_version_metadata()}, if_not_exists="create", stream_mode=list(_DASHBOARD_STREAM_MODES), @@ -357,7 +412,8 @@ async def create_dashboard_thread(login: str, body: ThreadCreateBody) -> dict[st login=login, repo_config=repo_config, repo_explicitly_none=body.repo_explicitly_none, - prompt=body.prompt.strip(), + prompt=body.prompt, + images=body.images, model_id=body.model_id, effort=body.effort, ) @@ -377,6 +433,7 @@ async def send_dashboard_message( owner, name, _ = _metadata_repo(metadata) prompt = body.content.strip() + content = _user_message_content(prompt, body.images) now_ms = _now_ms() chosen_model, chosen_effort = _normalize_model_choice(body.model_id, body.effort) metadata_update: dict[str, Any] = {"source": _DASHBOARD_SOURCE, "updated_at_ms": now_ms} @@ -386,9 +443,16 @@ async def send_dashboard_message( await client.threads.update(thread_id=thread_id, metadata=metadata_update) if await is_thread_active(thread_id): + queue_payload: dict[str, Any] = {"text": prompt, "source": _DASHBOARD_SOURCE} + if isinstance(content, list): + queue_payload["images"] = [ + block + for block in content + if isinstance(block, dict) and block.get("type") != "text" + ] queued = await queue_message_for_thread( thread_id, - {"text": prompt, "source": _DASHBOARD_SOURCE}, + queue_payload, ) if not queued: raise HTTPException(502, "failed to queue follow-up message") @@ -415,7 +479,7 @@ async def send_dashboard_message( run = await client.runs.create( thread_id, _ASSISTANT_ID, - input={"messages": [{"role": "user", "content": prompt}]}, + input={"messages": [{"role": "user", "content": content}]}, config={"configurable": configurable, "metadata": _agent_version_metadata()}, stream_mode=list(_DASHBOARD_STREAM_MODES), stream_resumable=True, diff --git a/agent/middleware/check_message_queue.py b/agent/middleware/check_message_queue.py index babdaeb7..e02554ab 100644 --- a/agent/middleware/check_message_queue.py +++ b/agent/middleware/check_message_queue.py @@ -39,9 +39,12 @@ async def _build_blocks_from_payload( ) -> list[dict[str, Any]]: text = payload.get("text", "") image_urls = payload.get("image_urls", []) or [] + images = payload.get("images", []) or [] blocks: list[dict[str, Any]] = [] if text: blocks.append({"type": "text", "text": text}) + if isinstance(images, list): + blocks.extend(image for image in images if isinstance(image, dict)) if not image_urls: return blocks @@ -119,7 +122,9 @@ async def check_message_queue_before_model( # noqa: PLR0911 content = msg.get("content") if _is_dashboard_queued_message(content): content_blocks.append({"type": "text", "text": DASHBOARD_HANDOFF_INSTRUCTION}) - if isinstance(content, dict) and ("text" in content or "image_urls" in content): + if isinstance(content, dict) and ( + "text" in content or "image_urls" in content or "images" in content + ): logger.debug("Queued message contains text + image URLs") blocks = await _build_blocks_from_payload(content) content_blocks.extend(blocks) diff --git a/tests/test_dashboard_message_adapter.py b/tests/test_dashboard_message_adapter.py index 415f437f..5de5babb 100644 --- a/tests/test_dashboard_message_adapter.py +++ b/tests/test_dashboard_message_adapter.py @@ -39,6 +39,38 @@ def test_state_messages_to_ui_maps_user_and_tool_calls() -> None: assert tool_chunk["output"] == "print('hi')" +def test_state_messages_to_ui_maps_user_images() -> None: + messages = [ + { + "type": "human", + "id": "u1", + "content": [ + { + "type": "image", + "source_type": "base64", + "data": "aW1hZ2U=", + "mime_type": "image/png", + }, + {"type": "text", "text": "What changed?"}, + ], + } + ] + + ui = state_messages_to_ui(messages) + + assert ui == [ + { + "id": "u1", + "author": "user", + "timestamp": ui[0]["timestamp"], + "chunks": [ + {"kind": "image", "base64": "aW1hZ2U=", "mimeType": "image/png"}, + {"kind": "text", "text": "What changed?"}, + ], + } + ] + + def test_state_messages_to_ui_tags_slack_and_linear_replies() -> None: messages = [ {"type": "human", "id": "u1", "content": "ping"}, diff --git a/tests/test_dashboard_web_handoff.py b/tests/test_dashboard_web_handoff.py index 4fd13b4a..bb17dc4a 100644 --- a/tests/test_dashboard_web_handoff.py +++ b/tests/test_dashboard_web_handoff.py @@ -91,6 +91,51 @@ async def test_dashboard_followup_on_slack_thread_uses_dashboard_source( assert run_config["repo"] == {"owner": "octo", "name": "repo"} +@pytest.mark.asyncio +async def test_dashboard_followup_sends_image_content_blocks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + metadata = { + "source": "dashboard", + "github_login": "octocat", + "repo_owner": "octo", + "repo_name": "repo", + } + client = _FakeClient(metadata) + + monkeypatch.setattr(thread_api, "langgraph_client", lambda: client) + monkeypatch.setattr(thread_api, "is_thread_active", _inactive_thread) + monkeypatch.setattr(thread_api, "_ensure_dashboard_github_token", _noop_token_check) + monkeypatch.setattr(thread_api, "get_profile", _empty_profile) + monkeypatch.setattr(thread_api, "_resolve_run_email", _run_email) + monkeypatch.setattr( + thread_api, + "create_image_block", + lambda *, base64, mime_type: {"type": "image", "data": base64, "mime_type": mime_type}, + ) + + await thread_api.send_dashboard_message( + "thread-1", + "octocat", + thread_api.ThreadMessageBody( + content="describe this", + images=[ + thread_api.DashboardImageBody( + base64="aW1hZ2U=", + mimeType="image/png", + fileName="screenshot.png", + ) + ], + ), + ) + + content = client.runs.created[0]["kwargs"]["input"]["messages"][0]["content"] + assert content == [ + {"type": "image", "data": "aW1hZ2U=", "mime_type": "image/png"}, + {"type": "text", "text": "describe this"}, + ] + + @pytest.mark.asyncio async def test_dashboard_followup_on_busy_thread_queues_dashboard_handoff( monkeypatch: pytest.MonkeyPatch, @@ -122,6 +167,48 @@ async def test_dashboard_followup_on_busy_thread_queues_dashboard_handoff( assert queued_messages == [{"text": "continue in web", "source": "dashboard"}] +@pytest.mark.asyncio +async def test_dashboard_followup_on_busy_thread_queues_images( + monkeypatch: pytest.MonkeyPatch, +) -> None: + metadata = { + "source": "dashboard", + "github_login": "octocat", + } + client = _FakeClient(metadata) + queued_messages: list[object] = [] + + async def fake_queue_message_for_thread(thread_id: str, message_content: object) -> bool: + queued_messages.append(message_content) + return True + + monkeypatch.setattr(thread_api, "langgraph_client", lambda: client) + monkeypatch.setattr(thread_api, "is_thread_active", _active_thread) + monkeypatch.setattr(thread_api, "queue_message_for_thread", fake_queue_message_for_thread) + monkeypatch.setattr( + thread_api, + "create_image_block", + lambda *, base64, mime_type: {"type": "image", "data": base64, "mime_type": mime_type}, + ) + + await thread_api.send_dashboard_message( + "thread-1", + "octocat", + thread_api.ThreadMessageBody( + content="continue in web", + images=[thread_api.DashboardImageBody(base64="aW1hZ2U=", mimeType="image/png")], + ), + ) + + assert queued_messages == [ + { + "text": "continue in web", + "source": "dashboard", + "images": [{"type": "image", "data": "aW1hZ2U=", "mime_type": "image/png"}], + } + ] + + @pytest.mark.asyncio async def test_dashboard_followup_preserves_explicit_repo_less_thread( monkeypatch: pytest.MonkeyPatch, diff --git a/ui/src/components/agents/AgentThreadView.tsx b/ui/src/components/agents/AgentThreadView.tsx index c91e4a5c..089af4e1 100644 --- a/ui/src/components/agents/AgentThreadView.tsx +++ b/ui/src/components/agents/AgentThreadView.tsx @@ -17,6 +17,33 @@ interface AgentThreadViewProps { thread: AgentThread; } +function messageText(message: Message): string { + return message.chunks + .filter((chunk) => chunk.kind === "text") + .map((chunk) => chunk.text) + .join(""); +} + +function messageImageKey(message: Message): string { + return message.chunks + .filter((chunk) => chunk.kind === "image") + .map((chunk) => `${chunk.mimeType}:${chunk.base64}`) + .join("\u0000"); +} + +function pendingImageKey(entry: PendingPrompt): string { + return (entry.images ?? []) + .map((image) => `${image.mimeType}:${image.base64}`) + .join("\u0000"); +} + +function isPendingPromptConfirmed(entry: PendingPrompt, messages: Array): boolean { + return messages.slice(entry.insertAt).some((message) => { + if (message.author !== "user") return false; + return messageText(message) === entry.prompt && messageImageKey(message) === pendingImageKey(entry); + }); +} + export function AgentThreadView({ user, thread }: AgentThreadViewProps) { const sendMessage = useSendAgentMessage(thread.id); const cancelThread = useCancelAgentThread(thread.id); @@ -37,39 +64,28 @@ export function AgentThreadView({ user, thread }: AgentThreadViewProps) { const [selection, setSelection] = useState(null); const activeSelection = selection ?? threadSelection ?? defaultSelection; - const userMessageTexts = useMemo(() => { - return new Set( - thread.messages - .filter((m) => m.author === "user") - .map((m) => - m.chunks - .filter((c) => c.kind === "text") - .map((c) => c.text) - .join(""), - ), - ); - }, [thread.messages]); - useEffect(() => { setPendingPrompts((prev) => { if (prev.length === 0) return prev; const next = dropPendingPrompts(thread.id, (entry) => - userMessageTexts.has(entry.prompt), + isPendingPromptConfirmed(entry, thread.messages), ); return next.length === prev.length ? prev : next; }); - }, [thread.id, userMessageTexts]); + }, [thread.id, thread.messages]); const displayMessages = useMemo>(() => { if (pendingPrompts.length === 0) return thread.messages; const baseTimestamp = new Date().toISOString(); const result = thread.messages.slice(); pendingPrompts.forEach((entry, i) => { + const chunks: Message["chunks"] = [...(entry.images ?? [])]; + if (entry.prompt) chunks.push({ kind: "text", text: entry.prompt }); const synth: Message = { id: `pending-user-${i}`, author: "user", timestamp: baseTimestamp, - chunks: [{ kind: "text", text: entry.prompt }], + chunks, }; const at = Math.min(Math.max(entry.insertAt, 0), result.length); result.splice(at, 0, synth); @@ -99,9 +115,10 @@ export function AgentThreadView({ user, thread }: AgentThreadViewProps) { compact busy={hasActiveRun} disabled={sendMessage.isPending} - onSubmit={(content) => + onSubmit={(content, images) => sendMessage.mutate({ content, + images, model_id: activeSelection?.modelId ?? null, effort: activeSelection?.effort ?? null, }) @@ -124,9 +141,10 @@ export function AgentThreadView({ user, thread }: AgentThreadViewProps) { compact busy={hasActiveRun} disabled={sendMessage.isPending} - onSubmit={(content) => + onSubmit={(content, images) => sendMessage.mutate({ content, + images, model_id: activeSelection?.modelId ?? null, effort: activeSelection?.effort ?? null, }) diff --git a/ui/src/components/agents/AgentsHome.tsx b/ui/src/components/agents/AgentsHome.tsx index 7067d66c..d7bd410e 100644 --- a/ui/src/components/agents/AgentsHome.tsx +++ b/ui/src/components/agents/AgentsHome.tsx @@ -32,9 +32,10 @@ export function AgentsHome() {
+ onSubmit={(prompt, images) => createThread.mutate({ prompt, + images, repo, repo_explicitly_none: repoOverride === null, model_id: activeSelection?.modelId ?? null, diff --git a/ui/src/components/agents/ported/CloudPromptBar.tsx b/ui/src/components/agents/ported/CloudPromptBar.tsx index c480afb3..1598550c 100644 --- a/ui/src/components/agents/ported/CloudPromptBar.tsx +++ b/ui/src/components/agents/ported/CloudPromptBar.tsx @@ -1,4 +1,11 @@ -import { ArrowUp, ChevronDown, LoaderCircle, Square } from "lucide-react" +import { + ArrowUp, + ChevronDown, + ImagePlus, + LoaderCircle, + Square, + X, +} from "lucide-react" import { memo, useCallback, @@ -10,19 +17,28 @@ import { } from "react" import type { ModelOption } from "@/lib/api" +import type { ImageChunk } from "@/lib/agents/types" import type { ModelSelection } from "@/lib/agents/useModelOptions" import { RepoSelector } from "@/components/agents/RepoSelector" import { formatModelSelection } from "@/lib/agents/useModelOptions" import { cn } from "@/lib/utils" const PROMPT_TEXTAREA_MAX_HEIGHT = 200 +const MAX_IMAGE_COUNT = 5 +const MAX_IMAGE_BYTES = 10 * 1024 * 1024 +const SUPPORTED_IMAGE_TYPES = new Set([ + "image/png", + "image/jpeg", + "image/gif", + "image/webp", +]) export interface CloudPromptBarProps { placeholder?: string compact?: boolean disabled?: boolean busy?: boolean - onSubmit?: (value: string) => void + onSubmit?: (value: string, images: Array) => void /** Called to stop the running agent. When set, the send button becomes a stop button while busy and the input is empty. */ onStop?: () => void stopping?: boolean @@ -35,6 +51,32 @@ export interface CloudPromptBarProps { onRepoChange?: (repo: string | null) => void } +function fileToImageChunk(file: File): Promise { + if (!SUPPORTED_IMAGE_TYPES.has(file.type) || file.size > MAX_IMAGE_BYTES) { + return Promise.resolve(null) + } + + return new Promise((resolve) => { + const reader = new FileReader() + reader.onload = () => { + const dataUrl = typeof reader.result === "string" ? reader.result : "" + const base64 = dataUrl.split(",")[1] + resolve( + base64 + ? { + kind: "image", + base64, + mimeType: file.type, + fileName: file.name, + } + : null + ) + } + reader.onerror = () => resolve(null) + reader.readAsDataURL(file) + }) +} + /** Web-adapted PromptBar from open-swe-app — local state, no Electron/Zustand deps. */ export const CloudPromptBar = memo(function CloudPromptBarComponent({ placeholder = "Ask Open SWE to build, fix bugs, explore", @@ -52,8 +94,12 @@ export const CloudPromptBar = memo(function CloudPromptBarComponent({ onRepoChange, }: CloudPromptBarProps) { const [value, setValue] = useState("") + const [pendingImages, setPendingImages] = useState>([]) + const [isDragOver, setIsDragOver] = useState(false) const [modelDropdownOpen, setModelDropdownOpen] = useState(false) const inputRef = useRef(null) + const fileInputRef = useRef(null) + const dragDepthRef = useRef(0) const modelDropdownRef = useRef(null) const combos = useMemo>(() => { @@ -68,12 +114,16 @@ export const CloudPromptBar = memo(function CloudPromptBarComponent({ const selectionLabel = formatModelSelection(models, selection) + const canSubmit = + !disabled && (value.trim().length > 0 || pendingImages.length > 0) + const handleSubmit = useCallback(() => { const trimmed = value.trim() - if (!trimmed || disabled) return - onSubmit?.(trimmed) + if (!canSubmit) return + onSubmit?.(trimmed, pendingImages) setValue("") - }, [disabled, onSubmit, value]) + setPendingImages([]) + }, [canSubmit, onSubmit, pendingImages, value]) useLayoutEffect(() => { const el = inputRef.current @@ -100,15 +150,69 @@ export const CloudPromptBar = memo(function CloudPromptBarComponent({ return () => document.removeEventListener("mousedown", handleClickOutside) }, []) + const addFiles = useCallback(async (files: FileList | Array) => { + const nextImages = await Promise.all( + Array.from(files).map(fileToImageChunk) + ) + const validImages = nextImages.filter( + (image): image is ImageChunk => image !== null + ) + if (validImages.length === 0) return + setPendingImages((prev) => + [...prev, ...validImages].slice(0, MAX_IMAGE_COUNT) + ) + }, []) + + const handleFileChange = useCallback( + (e: React.ChangeEvent) => { + const files = e.target.files + if (files) void addFiles(files) + e.target.value = "" + }, + [addFiles] + ) + + const handleDragEnter = useCallback((e: React.DragEvent) => { + if (!e.dataTransfer.types.includes("Files")) return + e.preventDefault() + dragDepthRef.current += 1 + setIsDragOver(true) + }, []) + + const handleDragOver = useCallback((e: React.DragEvent) => { + if (!e.dataTransfer.types.includes("Files")) return + e.preventDefault() + setIsDragOver(true) + }, []) + + const handleDragLeave = useCallback((e: React.DragEvent) => { + if (!e.dataTransfer.types.includes("Files")) return + e.preventDefault() + dragDepthRef.current = Math.max(0, dragDepthRef.current - 1) + if (dragDepthRef.current === 0) setIsDragOver(false) + }, []) + + const handleDrop = useCallback( + (e: React.DragEvent) => { + if (!e.dataTransfer.types.includes("Files")) return + e.preventDefault() + dragDepthRef.current = 0 + setIsDragOver(false) + void addFiles(e.dataTransfer.files) + }, + [addFiles] + ) + const handleKeyDown = (e: React.KeyboardEvent) => { - if (e.key === "Enter" && !e.shiftKey && value.trim()) { + if (e.key === "Enter" && !e.shiftKey && canSubmit) { e.preventDefault() handleSubmit() } } const pickerDisabled = combos.length === 0 || !onSelectionChange - const showStop = busy && !disabled && !value.trim() && !!onStop + const showStop = + busy && !disabled && !value.trim() && pendingImages.length === 0 && !!onStop return (
)}
+ {isDragOver && ( +
+ + Drop images here + +
+ )} + + + + {pendingImages.length > 0 && ( +
+ {pendingImages.map((image, index) => ( +
+ {image.fileName + +
+ ))} +
+ )} +