feat: support dashboard chat images (#1435)

* feat: support dashboard chat images

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>

* fix: keep pending image prompts visible

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>

---------

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
Johannes du Plessis 2026-06-06 11:50:24 -07:00 • committed by GitHub
parent 1f619eae0f
commit 072c0158ff
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 503 additions and 52 deletions

View file

@ -3,6 +3,7 @@
from __future__ import annotations from __future__ import annotations
import json import json
import re
import uuid import uuid
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import Any 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"}) _EXECUTE_TOOLS = frozenset({"execute", "bash", "shell", "run_terminal_cmd"})
_SEARCH_TOOLS = frozenset({"glob", "grep", "web_search", "fetch_url", "search"}) _SEARCH_TOOLS = frozenset({"glob", "grep", "web_search", "fetch_url", "search"})
_INTERNAL_TOOLS = frozenset({"confirming_completion", "no_op"}) _INTERNAL_TOOLS = frozenset({"confirming_completion", "no_op"})
_DATA_IMAGE_RE = re.compile(r"^data:(image/[^;]+);base64,(.+)$", re.DOTALL)
def _now_iso() -> str: def _now_iso() -> str:
@ -78,6 +80,43 @@ def _parse_tool_args(raw: Any) -> dict[str, Any]:
return {} 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: 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") path = args.get("path") or args.get("file_path") or args.get("target_file")
if not isinstance(path, str) or not path.strip(): 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"]) agent_turn["chunks"] = _merge_text_chunks(agent_turn["chunks"])
ui_messages.append(agent_turn) ui_messages.append(agent_turn)
agent_turn = None agent_turn = None
text = extract_text_content(raw.get("content", "")) content = raw.get("content", "")
if not text: chunks = _image_chunks(content)
text = extract_text_content(content)
if text:
chunks.append({"kind": "text", "text": text})
if not chunks:
continue continue
ui_messages.append( ui_messages.append(
{ {
"id": msg_id, "id": msg_id,
"author": "user", "author": "user",
"timestamp": timestamp, "timestamp": timestamp,
"chunks": [{"kind": "text", "text": text}], "chunks": chunks,
} }
) )
continue continue

View file

@ -2,6 +2,8 @@
from __future__ import annotations from __future__ import annotations
import base64
import binascii
import json import json
import logging import logging
import os import os
@ -11,8 +13,9 @@ from datetime import UTC, datetime
from typing import Any from typing import Any
from fastapi import HTTPException from fastapi import HTTPException
from langchain_core.messages.content import create_image_block
from langgraph_sdk.errors import InternalServerError 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 ..utils.thread_ops import is_thread_active, langgraph_client, queue_message_for_thread
from .message_adapter import state_messages_to_ui from .message_adapter import state_messages_to_ui
@ -25,6 +28,9 @@ logger = logging.getLogger(__name__)
_ASSISTANT_ID = "agent" _ASSISTANT_ID = "agent"
_DASHBOARD_SOURCE = "dashboard" _DASHBOARD_SOURCE = "dashboard"
_DASHBOARD_STREAM_MODES: tuple[str, ...] = ("values", "updates", "messages-tuple") _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"). # Sources whose threads should surface in the Agents UI (besides "dashboard").
_SURFACED_SOURCES: tuple[str, ...] = ("dashboard", "github", "slack", "linear", "schedule") _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") 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): 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: str | None = None
repo_explicitly_none: bool = False repo_explicitly_none: bool = False
model_id: str | None = None model_id: str | None = None
@ -54,7 +70,8 @@ class ThreadCreateBody(BaseModel):
class ThreadMessageBody(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 model_id: str | None = None
effort: 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} 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: async def _ensure_dashboard_github_token(login: str) -> None:
token = await get_valid_access_token(login) token = await get_valid_access_token(login)
if not token: if not token:
@ -283,12 +335,15 @@ async def _start_agent_run(
repo_config: dict[str, str], repo_config: dict[str, str],
repo_explicitly_none: bool = False, repo_explicitly_none: bool = False,
prompt: str, prompt: str,
images: list[DashboardImageBody] | None = None,
title: str | None = None, title: str | None = None,
model_id: str | None = None, model_id: str | None = None,
effort: str | None = None, effort: str | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
profile = await get_profile(login) or {} profile = await get_profile(login) or {}
now_ms = _now_ms() now_ms = _now_ms()
prompt = prompt.strip()
content = _user_message_content(prompt, images or [])
chosen_model, chosen_effort = _normalize_model_choice(model_id, effort) chosen_model, chosen_effort = _normalize_model_choice(model_id, effort)
metadata_model = chosen_model or profile.get("default_model") or "Default" metadata_model = chosen_model or profile.get("default_model") or "Default"
metadata_effort = chosen_effort or profile.get("reasoning_effort") metadata_effort = chosen_effort or profile.get("reasoning_effort")
@ -332,7 +387,7 @@ async def _start_agent_run(
run = await client.runs.create( run = await client.runs.create(
thread_id, thread_id,
_ASSISTANT_ID, _ASSISTANT_ID,
input={"messages": [{"role": "user", "content": prompt}]}, input={"messages": [{"role": "user", "content": content}]},
config={"configurable": configurable, "metadata": _agent_version_metadata()}, config={"configurable": configurable, "metadata": _agent_version_metadata()},
if_not_exists="create", if_not_exists="create",
stream_mode=list(_DASHBOARD_STREAM_MODES), stream_mode=list(_DASHBOARD_STREAM_MODES),
@ -357,7 +412,8 @@ async def create_dashboard_thread(login: str, body: ThreadCreateBody) -> dict[st
login=login, login=login,
repo_config=repo_config, repo_config=repo_config,
repo_explicitly_none=body.repo_explicitly_none, repo_explicitly_none=body.repo_explicitly_none,
prompt=body.prompt.strip(), prompt=body.prompt,
images=body.images,
model_id=body.model_id, model_id=body.model_id,
effort=body.effort, effort=body.effort,
) )
@ -377,6 +433,7 @@ async def send_dashboard_message(
owner, name, _ = _metadata_repo(metadata) owner, name, _ = _metadata_repo(metadata)
prompt = body.content.strip() prompt = body.content.strip()
content = _user_message_content(prompt, body.images)
now_ms = _now_ms() now_ms = _now_ms()
chosen_model, chosen_effort = _normalize_model_choice(body.model_id, body.effort) 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} 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) await client.threads.update(thread_id=thread_id, metadata=metadata_update)
if await is_thread_active(thread_id): 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( queued = await queue_message_for_thread(
thread_id, thread_id,
{"text": prompt, "source": _DASHBOARD_SOURCE}, queue_payload,
) )
if not queued: if not queued:
raise HTTPException(502, "failed to queue follow-up message") raise HTTPException(502, "failed to queue follow-up message")
@ -415,7 +479,7 @@ async def send_dashboard_message(
run = await client.runs.create( run = await client.runs.create(
thread_id, thread_id,
_ASSISTANT_ID, _ASSISTANT_ID,
input={"messages": [{"role": "user", "content": prompt}]}, input={"messages": [{"role": "user", "content": content}]},
config={"configurable": configurable, "metadata": _agent_version_metadata()}, config={"configurable": configurable, "metadata": _agent_version_metadata()},
stream_mode=list(_DASHBOARD_STREAM_MODES), stream_mode=list(_DASHBOARD_STREAM_MODES),
stream_resumable=True, stream_resumable=True,

View file

@ -39,9 +39,12 @@ async def _build_blocks_from_payload(
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
text = payload.get("text", "") text = payload.get("text", "")
image_urls = payload.get("image_urls", []) or [] image_urls = payload.get("image_urls", []) or []
images = payload.get("images", []) or []
blocks: list[dict[str, Any]] = [] blocks: list[dict[str, Any]] = []
if text: if text:
blocks.append({"type": "text", "text": 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: if not image_urls:
return blocks return blocks
@ -119,7 +122,9 @@ async def check_message_queue_before_model( # noqa: PLR0911
content = msg.get("content") content = msg.get("content")
if _is_dashboard_queued_message(content): if _is_dashboard_queued_message(content):
content_blocks.append({"type": "text", "text": DASHBOARD_HANDOFF_INSTRUCTION}) 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") logger.debug("Queued message contains text + image URLs")
blocks = await _build_blocks_from_payload(content) blocks = await _build_blocks_from_payload(content)
content_blocks.extend(blocks) content_blocks.extend(blocks)

View file

@ -39,6 +39,38 @@ def test_state_messages_to_ui_maps_user_and_tool_calls() -> None:
assert tool_chunk["output"] == "print('hi')" 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: def test_state_messages_to_ui_tags_slack_and_linear_replies() -> None:
messages = [ messages = [
{"type": "human", "id": "u1", "content": "ping"}, {"type": "human", "id": "u1", "content": "ping"},

View file

@ -91,6 +91,51 @@ async def test_dashboard_followup_on_slack_thread_uses_dashboard_source(
assert run_config["repo"] == {"owner": "octo", "name": "repo"} 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 @pytest.mark.asyncio
async def test_dashboard_followup_on_busy_thread_queues_dashboard_handoff( async def test_dashboard_followup_on_busy_thread_queues_dashboard_handoff(
monkeypatch: pytest.MonkeyPatch, 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"}] 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 @pytest.mark.asyncio
async def test_dashboard_followup_preserves_explicit_repo_less_thread( async def test_dashboard_followup_preserves_explicit_repo_less_thread(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,

View file

@ -17,6 +17,33 @@ interface AgentThreadViewProps {
thread: AgentThread; 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<Message>): 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) { export function AgentThreadView({ user, thread }: AgentThreadViewProps) {
const sendMessage = useSendAgentMessage(thread.id); const sendMessage = useSendAgentMessage(thread.id);
const cancelThread = useCancelAgentThread(thread.id); const cancelThread = useCancelAgentThread(thread.id);
@ -37,39 +64,28 @@ export function AgentThreadView({ user, thread }: AgentThreadViewProps) {
const [selection, setSelection] = useState<ModelSelection | null>(null); const [selection, setSelection] = useState<ModelSelection | null>(null);
const activeSelection = selection ?? threadSelection ?? defaultSelection; 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(() => { useEffect(() => {
setPendingPrompts((prev) => { setPendingPrompts((prev) => {
if (prev.length === 0) return prev; if (prev.length === 0) return prev;
const next = dropPendingPrompts(thread.id, (entry) => const next = dropPendingPrompts(thread.id, (entry) =>
userMessageTexts.has(entry.prompt), isPendingPromptConfirmed(entry, thread.messages),
); );
return next.length === prev.length ? prev : next; return next.length === prev.length ? prev : next;
}); });
}, [thread.id, userMessageTexts]); }, [thread.id, thread.messages]);
const displayMessages = useMemo<Array<Message>>(() => { const displayMessages = useMemo<Array<Message>>(() => {
if (pendingPrompts.length === 0) return thread.messages; if (pendingPrompts.length === 0) return thread.messages;
const baseTimestamp = new Date().toISOString(); const baseTimestamp = new Date().toISOString();
const result = thread.messages.slice(); const result = thread.messages.slice();
pendingPrompts.forEach((entry, i) => { pendingPrompts.forEach((entry, i) => {
const chunks: Message["chunks"] = [...(entry.images ?? [])];
if (entry.prompt) chunks.push({ kind: "text", text: entry.prompt });
const synth: Message = { const synth: Message = {
id: `pending-user-${i}`, id: `pending-user-${i}`,
author: "user", author: "user",
timestamp: baseTimestamp, timestamp: baseTimestamp,
chunks: [{ kind: "text", text: entry.prompt }], chunks,
}; };
const at = Math.min(Math.max(entry.insertAt, 0), result.length); const at = Math.min(Math.max(entry.insertAt, 0), result.length);
result.splice(at, 0, synth); result.splice(at, 0, synth);
@ -99,9 +115,10 @@ export function AgentThreadView({ user, thread }: AgentThreadViewProps) {
compact compact
busy={hasActiveRun} busy={hasActiveRun}
disabled={sendMessage.isPending} disabled={sendMessage.isPending}
onSubmit={(content) => onSubmit={(content, images) =>
sendMessage.mutate({ sendMessage.mutate({
content, content,
images,
model_id: activeSelection?.modelId ?? null, model_id: activeSelection?.modelId ?? null,
effort: activeSelection?.effort ?? null, effort: activeSelection?.effort ?? null,
}) })
@ -124,9 +141,10 @@ export function AgentThreadView({ user, thread }: AgentThreadViewProps) {
compact compact
busy={hasActiveRun} busy={hasActiveRun}
disabled={sendMessage.isPending} disabled={sendMessage.isPending}
onSubmit={(content) => onSubmit={(content, images) =>
sendMessage.mutate({ sendMessage.mutate({
content, content,
images,
model_id: activeSelection?.modelId ?? null, model_id: activeSelection?.modelId ?? null,
effort: activeSelection?.effort ?? null, effort: activeSelection?.effort ?? null,
}) })

View file

@ -32,9 +32,10 @@ export function AgentsHome() {
<div className="flex w-full flex-col items-center gap-6"> <div className="flex w-full flex-col items-center gap-6">
<Logo /> <Logo />
<AgentPromptBar <AgentPromptBar
onSubmit={(prompt) => onSubmit={(prompt, images) =>
createThread.mutate({ createThread.mutate({
prompt, prompt,
images,
repo, repo,
repo_explicitly_none: repoOverride === null, repo_explicitly_none: repoOverride === null,
model_id: activeSelection?.modelId ?? null, model_id: activeSelection?.modelId ?? null,

View file

@ -1,4 +1,11 @@
import { ArrowUp, ChevronDown, LoaderCircle, Square } from "lucide-react" import {
ArrowUp,
ChevronDown,
ImagePlus,
LoaderCircle,
Square,
X,
} from "lucide-react"
import { import {
memo, memo,
useCallback, useCallback,
@ -10,19 +17,28 @@ import {
} from "react" } from "react"
import type { ModelOption } from "@/lib/api" import type { ModelOption } from "@/lib/api"
import type { ImageChunk } from "@/lib/agents/types"
import type { ModelSelection } from "@/lib/agents/useModelOptions" import type { ModelSelection } from "@/lib/agents/useModelOptions"
import { RepoSelector } from "@/components/agents/RepoSelector" import { RepoSelector } from "@/components/agents/RepoSelector"
import { formatModelSelection } from "@/lib/agents/useModelOptions" import { formatModelSelection } from "@/lib/agents/useModelOptions"
import { cn } from "@/lib/utils" import { cn } from "@/lib/utils"
const PROMPT_TEXTAREA_MAX_HEIGHT = 200 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 { export interface CloudPromptBarProps {
placeholder?: string placeholder?: string
compact?: boolean compact?: boolean
disabled?: boolean disabled?: boolean
busy?: boolean busy?: boolean
onSubmit?: (value: string) => void onSubmit?: (value: string, images: Array<ImageChunk>) => void
/** Called to stop the running agent. When set, the send button becomes a stop button while busy and the input is empty. */ /** Called to stop the running agent. When set, the send button becomes a stop button while busy and the input is empty. */
onStop?: () => void onStop?: () => void
stopping?: boolean stopping?: boolean
@ -35,6 +51,32 @@ export interface CloudPromptBarProps {
onRepoChange?: (repo: string | null) => void onRepoChange?: (repo: string | null) => void
} }
function fileToImageChunk(file: File): Promise<ImageChunk | null> {
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. */ /** Web-adapted PromptBar from open-swe-app — local state, no Electron/Zustand deps. */
export const CloudPromptBar = memo(function CloudPromptBarComponent({ export const CloudPromptBar = memo(function CloudPromptBarComponent({
placeholder = "Ask Open SWE to build, fix bugs, explore", placeholder = "Ask Open SWE to build, fix bugs, explore",
@ -52,8 +94,12 @@ export const CloudPromptBar = memo(function CloudPromptBarComponent({
onRepoChange, onRepoChange,
}: CloudPromptBarProps) { }: CloudPromptBarProps) {
const [value, setValue] = useState("") const [value, setValue] = useState("")
const [pendingImages, setPendingImages] = useState<Array<ImageChunk>>([])
const [isDragOver, setIsDragOver] = useState(false)
const [modelDropdownOpen, setModelDropdownOpen] = useState(false) const [modelDropdownOpen, setModelDropdownOpen] = useState(false)
const inputRef = useRef<HTMLTextAreaElement>(null) const inputRef = useRef<HTMLTextAreaElement>(null)
const fileInputRef = useRef<HTMLInputElement>(null)
const dragDepthRef = useRef(0)
const modelDropdownRef = useRef<HTMLDivElement>(null) const modelDropdownRef = useRef<HTMLDivElement>(null)
const combos = useMemo<Array<ModelSelection>>(() => { const combos = useMemo<Array<ModelSelection>>(() => {
@ -68,12 +114,16 @@ export const CloudPromptBar = memo(function CloudPromptBarComponent({
const selectionLabel = formatModelSelection(models, selection) const selectionLabel = formatModelSelection(models, selection)
const canSubmit =
!disabled && (value.trim().length > 0 || pendingImages.length > 0)
const handleSubmit = useCallback(() => { const handleSubmit = useCallback(() => {
const trimmed = value.trim() const trimmed = value.trim()
if (!trimmed || disabled) return if (!canSubmit) return
onSubmit?.(trimmed) onSubmit?.(trimmed, pendingImages)
setValue("") setValue("")
}, [disabled, onSubmit, value]) setPendingImages([])
}, [canSubmit, onSubmit, pendingImages, value])
useLayoutEffect(() => { useLayoutEffect(() => {
const el = inputRef.current const el = inputRef.current
@ -100,15 +150,69 @@ export const CloudPromptBar = memo(function CloudPromptBarComponent({
return () => document.removeEventListener("mousedown", handleClickOutside) return () => document.removeEventListener("mousedown", handleClickOutside)
}, []) }, [])
const addFiles = useCallback(async (files: FileList | Array<File>) => {
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<HTMLInputElement>) => {
const files = e.target.files
if (files) void addFiles(files)
e.target.value = ""
},
[addFiles]
)
const handleDragEnter = useCallback((e: React.DragEvent<HTMLDivElement>) => {
if (!e.dataTransfer.types.includes("Files")) return
e.preventDefault()
dragDepthRef.current += 1
setIsDragOver(true)
}, [])
const handleDragOver = useCallback((e: React.DragEvent<HTMLDivElement>) => {
if (!e.dataTransfer.types.includes("Files")) return
e.preventDefault()
setIsDragOver(true)
}, [])
const handleDragLeave = useCallback((e: React.DragEvent<HTMLDivElement>) => {
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<HTMLDivElement>) => {
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<HTMLTextAreaElement>) => { const handleKeyDown = (e: React.KeyboardEvent<HTMLTextAreaElement>) => {
if (e.key === "Enter" && !e.shiftKey && value.trim()) { if (e.key === "Enter" && !e.shiftKey && canSubmit) {
e.preventDefault() e.preventDefault()
handleSubmit() handleSubmit()
} }
} }
const pickerDisabled = combos.length === 0 || !onSelectionChange const pickerDisabled = combos.length === 0 || !onSelectionChange
const showStop = busy && !disabled && !value.trim() && !!onStop const showStop =
busy && !disabled && !value.trim() && pendingImages.length === 0 && !!onStop
return ( return (
<div <div
@ -127,11 +231,62 @@ export const CloudPromptBar = memo(function CloudPromptBarComponent({
</div> </div>
)} )}
<div <div
onDragEnter={handleDragEnter}
onDragOver={handleDragOver}
onDragLeave={handleDragLeave}
onDrop={handleDrop}
className={cn( className={cn(
"relative flex min-h-[106px] flex-col rounded-2xl border border-[var(--ui-border)] bg-[var(--ui-surface)] px-3 py-2.5 shadow-sm", "relative flex min-h-[106px] flex-col rounded-2xl border border-[var(--ui-border)] bg-[var(--ui-surface)] px-3 py-2.5 shadow-sm",
compact && "min-h-[88px]" compact && "min-h-[88px]",
isDragOver && "border-[var(--ui-accent)]"
)} )}
> >
{isDragOver && (
<div className="pointer-events-none absolute inset-0 z-20 flex items-center justify-center rounded-2xl bg-[var(--ui-surface)]/80 backdrop-blur-sm">
<span className="rounded-md bg-[var(--ui-panel-2)] px-3 py-1.5 text-sm font-medium text-[color:var(--ui-accent)]">
Drop images here
</span>
</div>
)}
<input
ref={fileInputRef}
type="file"
accept="image/png,image/jpeg,image/gif,image/webp"
multiple
className="hidden"
onChange={handleFileChange}
/>
{pendingImages.length > 0 && (
<div className="mb-2 flex flex-wrap gap-2">
{pendingImages.map((image, index) => (
<div
key={`${image.fileName ?? "image"}-${index}`}
className="group relative"
>
<img
src={`data:${image.mimeType};base64,${image.base64}`}
alt={image.fileName || "pending image"}
className="size-16 rounded-lg border border-[var(--ui-border)] object-cover"
/>
<button
type="button"
aria-label="Remove image"
onClick={() =>
setPendingImages((prev) =>
prev.filter((_, i) => i !== index)
)
}
className="absolute -top-1.5 -right-1.5 flex size-5 items-center justify-center rounded-full border border-[var(--ui-border)] bg-[var(--ui-panel-2)] text-[color:var(--ui-text-muted)] opacity-0 shadow-sm transition-opacity group-hover:opacity-100 hover:text-[color:var(--ui-text)]"
>
<X className="size-3" />
</button>
</div>
))}
</div>
)}
<textarea <textarea
ref={inputRef} ref={inputRef}
rows={1} rows={1}
@ -195,13 +350,23 @@ export const CloudPromptBar = memo(function CloudPromptBarComponent({
)} )}
</div> </div>
<button
type="button"
onClick={() => fileInputRef.current?.click()}
disabled={disabled || pendingImages.length >= MAX_IMAGE_COUNT}
aria-label="Attach images"
className="ml-auto flex size-7 shrink-0 items-center justify-center rounded-full text-[color:var(--ui-text-muted)] transition-colors hover:bg-[var(--ui-panel-2)] hover:text-[color:var(--ui-text)] disabled:cursor-default disabled:opacity-40"
>
<ImagePlus className="size-4" />
</button>
{showStop ? ( {showStop ? (
<button <button
type="button" type="button"
onClick={onStop} onClick={onStop}
disabled={stopping} disabled={stopping}
aria-label="Stop run" aria-label="Stop run"
className="ml-auto flex size-7 shrink-0 items-center justify-center rounded-full bg-[var(--ui-accent)] text-white transition-opacity hover:opacity-90 disabled:cursor-default disabled:opacity-40" className="flex size-7 shrink-0 items-center justify-center rounded-full bg-[var(--ui-accent)] text-white transition-opacity hover:opacity-90 disabled:cursor-default disabled:opacity-40"
> >
{stopping ? ( {stopping ? (
<LoaderCircle className="size-3.5 animate-spin" /> <LoaderCircle className="size-3.5 animate-spin" />
@ -213,9 +378,9 @@ export const CloudPromptBar = memo(function CloudPromptBarComponent({
<button <button
type="button" type="button"
onClick={handleSubmit} onClick={handleSubmit}
disabled={!value.trim() || disabled} disabled={!canSubmit}
aria-label="Send message" aria-label="Send message"
className="ml-auto flex size-7 shrink-0 items-center justify-center rounded-full bg-[var(--ui-accent)] text-white transition-opacity hover:opacity-90 disabled:cursor-default disabled:opacity-40" className="flex size-7 shrink-0 items-center justify-center rounded-full bg-[var(--ui-accent)] text-white transition-opacity hover:opacity-90 disabled:cursor-default disabled:opacity-40"
> >
{disabled ? ( {disabled ? (
<LoaderCircle className="size-3.5 animate-spin" /> <LoaderCircle className="size-3.5 animate-spin" />

View file

@ -1,9 +1,10 @@
import type { AgentSchedule, AgentThread, Message } from "./types" import type { AgentSchedule, AgentThread, ImageChunk, Message } from "./types"
export type { AgentSchedule, AgentThread, Message } export type { AgentSchedule, AgentThread, Message }
export interface ThreadCreateRequest { export interface ThreadCreateRequest {
prompt: string prompt: string
images?: Array<ImageChunk>
repo?: string | null repo?: string | null
repo_explicitly_none?: boolean repo_explicitly_none?: boolean
model_id?: string | null model_id?: string | null
@ -12,6 +13,7 @@ export interface ThreadCreateRequest {
export interface ThreadMessageRequest { export interface ThreadMessageRequest {
content: string content: string
images?: Array<ImageChunk>
model_id?: string | null model_id?: string | null
effort?: string | null effort?: string | null
} }

View file

@ -1,8 +1,22 @@
import type { ImageChunk } from "./types";
const STORAGE_KEY = (threadId: string) => `open-swe:pending-prompts:${threadId}`; const STORAGE_KEY = (threadId: string) => `open-swe:pending-prompts:${threadId}`;
export interface PendingPrompt { export interface PendingPrompt {
prompt: string; prompt: string;
insertAt: number; insertAt: number;
images?: Array<ImageChunk>;
}
function isImageChunk(value: unknown): value is ImageChunk {
const image = value as { base64?: unknown; kind?: unknown; mimeType?: unknown };
return (
typeof value === "object" &&
value !== null &&
image.kind === "image" &&
typeof image.base64 === "string" &&
typeof image.mimeType === "string"
);
} }
function isPendingPrompt(value: unknown): value is PendingPrompt { function isPendingPrompt(value: unknown): value is PendingPrompt {
@ -10,11 +24,14 @@ function isPendingPrompt(value: unknown): value is PendingPrompt {
typeof value === "object" && typeof value === "object" &&
value !== null && value !== null &&
typeof (value as PendingPrompt).prompt === "string" && typeof (value as PendingPrompt).prompt === "string" &&
typeof (value as PendingPrompt).insertAt === "number" typeof (value as PendingPrompt).insertAt === "number" &&
((value as PendingPrompt).images === undefined ||
(Array.isArray((value as PendingPrompt).images) &&
(value as PendingPrompt).images!.every(isImageChunk)))
); );
} }
function safeRead(threadId: string): PendingPrompt[] { function safeRead(threadId: string): Array<PendingPrompt> {
if (typeof window === "undefined") return []; if (typeof window === "undefined") return [];
try { try {
const raw = window.sessionStorage.getItem(STORAGE_KEY(threadId)); const raw = window.sessionStorage.getItem(STORAGE_KEY(threadId));
@ -26,27 +43,36 @@ function safeRead(threadId: string): PendingPrompt[] {
} }
} }
function safeWrite(threadId: string, prompts: PendingPrompt[]): void { function safeWrite(threadId: string, prompts: Array<PendingPrompt>): void {
if (typeof window === "undefined") return; if (typeof window === "undefined") return;
if (prompts.length === 0) { if (prompts.length === 0) {
window.sessionStorage.removeItem(STORAGE_KEY(threadId)); window.sessionStorage.removeItem(STORAGE_KEY(threadId));
return; return;
} }
window.sessionStorage.setItem(STORAGE_KEY(threadId), JSON.stringify(prompts)); try {
window.sessionStorage.setItem(STORAGE_KEY(threadId), JSON.stringify(prompts));
} catch {
window.sessionStorage.removeItem(STORAGE_KEY(threadId));
}
} }
export function getPendingPrompts(threadId: string): PendingPrompt[] { export function getPendingPrompts(threadId: string): Array<PendingPrompt> {
return safeRead(threadId); return safeRead(threadId);
} }
export function addPendingPrompt(threadId: string, prompt: string, insertAt: number): void { export function addPendingPrompt(
safeWrite(threadId, [...safeRead(threadId), { prompt, insertAt }]); threadId: string,
prompt: string,
insertAt: number,
images?: Array<ImageChunk>,
): void {
safeWrite(threadId, [...safeRead(threadId), { prompt, insertAt, images }]);
} }
export function dropPendingPrompts( export function dropPendingPrompts(
threadId: string, threadId: string,
predicate: (entry: PendingPrompt) => boolean, predicate: (entry: PendingPrompt) => boolean,
): PendingPrompt[] { ): Array<PendingPrompt> {
const next = safeRead(threadId).filter((p) => !predicate(p)); const next = safeRead(threadId).filter((p) => !predicate(p));
safeWrite(threadId, next); safeWrite(threadId, next);
return next; return next;

View file

@ -4,6 +4,7 @@ import { useNavigate } from "@tanstack/react-router"
import { agentsApi } from "./api" import { agentsApi } from "./api"
import { addPendingPrompt } from "./pendingPrompts" import { addPendingPrompt } from "./pendingPrompts"
import type { ScheduleUpdateRequest } from "./api" import type { ScheduleUpdateRequest } from "./api"
import type { ImageChunk } from "./types"
export const agentThreadKeys = { export const agentThreadKeys = {
all: ["agent-threads"] as const, all: ["agent-threads"] as const,
@ -80,7 +81,12 @@ export function useCreateAgentThread() {
return useMutation({ return useMutation({
mutationFn: agentsApi.createThread, mutationFn: agentsApi.createThread,
onSuccess: (thread, variables) => { onSuccess: (thread, variables) => {
addPendingPrompt(thread.id, variables.prompt, thread.messages.length) addPendingPrompt(
thread.id,
variables.prompt,
thread.messages.length,
variables.images
)
queryClient.setQueryData(agentThreadKeys.detail(thread.id), { queryClient.setQueryData(agentThreadKeys.detail(thread.id), {
...thread, ...thread,
status: thread.status === "idle" ? "running" : thread.status, status: thread.status === "idle" ? "running" : thread.status,
@ -93,6 +99,7 @@ export function useCreateAgentThread() {
export interface SendAgentMessageVariables { export interface SendAgentMessageVariables {
content: string content: string
images?: Array<ImageChunk>
model_id?: string | null model_id?: string | null
effort?: string | null effort?: string | null
} }
@ -104,6 +111,7 @@ export function useSendAgentMessage(threadId: string) {
mutationFn: (vars: SendAgentMessageVariables) => mutationFn: (vars: SendAgentMessageVariables) =>
agentsApi.sendMessage(threadId, { agentsApi.sendMessage(threadId, {
content: vars.content, content: vars.content,
images: vars.images,
model_id: vars.model_id, model_id: vars.model_id,
effort: vars.effort, effort: vars.effort,
}), }),
@ -114,7 +122,7 @@ export function useSendAgentMessage(threadId: string) {
const insertAt = Array.isArray(cached?.messages) const insertAt = Array.isArray(cached?.messages)
? cached.messages.length ? cached.messages.length
: 0 : 0
addPendingPrompt(threadId, vars.content, insertAt) addPendingPrompt(threadId, vars.content, insertAt, vars.images)
}, },
onSuccess: (thread) => { onSuccess: (thread) => {
queryClient.setQueryData( queryClient.setQueryData(