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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions web/components/ai-elements/prompt-input.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -927,6 +927,15 @@ export const PromptInput = ({
form.reset();
}

const restoreText = () => {
if (usingProvider) return;
const input = form.elements.namedItem("message");
if (input instanceof HTMLTextAreaElement && input.value === "") {
input.value = text;
input.dispatchEvent(new Event("input", { bubbles: true }));
}
};

try {
// Convert blob URLs to data URLs asynchronously
const convertedFiles: FileUIPart[] = await Promise.all(
Expand Down Expand Up @@ -955,6 +964,7 @@ export const PromptInput = ({
}
} catch {
// Don't clear on error - user may want to retry
restoreText();
}
} else {
// Sync function completed without throwing, clear inputs
Expand All @@ -965,6 +975,7 @@ export const PromptInput = ({
}
} catch {
// Don't clear on error - user may want to retry
restoreText();
}
},
[usingProvider, controller, files, onSubmit, clear]
Expand Down Expand Up @@ -1086,8 +1097,10 @@ export const PromptInputTextarea = ({
});
};
form.addEventListener("reset", handleReset);
textarea.addEventListener("input", handleReset);
return () => {
form.removeEventListener("reset", handleReset);
textarea.removeEventListener("input", handleReset);
};
}, []);

Expand Down
118 changes: 117 additions & 1 deletion web/components/ai-elements/tests/prompt-input.test.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,25 @@ import {
type ComponentType,
type ReactNode,
} from "react";
import type * as JSXDevelopmentRuntime from "react/jsx-dev-runtime";
import { z } from "zod";
import { renderToStaticMarkup } from "react-dom/server";
import { describe, expect, it, vi } from "vitest";
import { afterEach, describe, expect, it, vi } from "vitest";

const captured = vi.hoisted(() => ({
forms: new Array<unknown>(),
}));

vi.mock("react/jsx-dev-runtime", async (importOriginal) => {
const runtime = await importOriginal<typeof JSXDevelopmentRuntime>();
return {
...runtime,
jsxDEV: (...args: Parameters<typeof runtime.jsxDEV>) => {
if (args[0] === "form") captured.forms.push(args[1]);
return runtime.jsxDEV(...args);
},
};
});

vi.mock("motion/react", () => {
// oxlint-disable-next-line unicorn/consistent-function-scoping -- Vitest hoists mock factories above module-scope component values.
Expand Down Expand Up @@ -68,6 +85,11 @@ import {
} from "@web/components/ai-elements/prompt-input";

describe("prompt input", () => {
afterEach(() => {
vi.unstubAllGlobals();
captured.forms.length = 0;
});

it("anchors the compact submit button without dropping footer children", () => {
const markup = renderToStaticMarkup(
<PromptInput compact onSubmit={() => undefined}>
Expand Down Expand Up @@ -109,4 +131,98 @@ describe("prompt input", () => {
);
expect(regularMarkup).not.toContain('data-motion-element="div"');
});

it.each(["async", "sync"] as const)(
"restores the submitted text after %s send failure",
async (mode) => {
const onSubmit =
mode === "async"
? () => Promise.reject(new Error("Send unavailable"))
: () => {
throw new Error("Send unavailable");
};
const input = submitDraft(onSubmit, "Keep my travel notes");

await vi.waitFor(() => {
expect(input.value).toBe("Keep my travel notes");
});
expect(input.dispatchEvent).toHaveBeenCalledWith(
expect.objectContaining({ bubbles: true, type: "input" })
);
}
);

it("leaves successfully submitted text cleared", async () => {
const onSubmit = vi.fn<() => Promise<void>>(() => Promise.resolve());
const input = submitDraft(onSubmit, "Send my travel notes");

await vi.waitFor(() => {
expect(onSubmit).toHaveBeenCalledWith(
{ files: [], text: "Send my travel notes" },
expect.anything()
);
});
expect(input.value).toBe("");
expect(input.dispatchEvent).not.toHaveBeenCalled();
});

it("keeps a newer draft when the previous send fails", async () => {
const pending = Promise.withResolvers<undefined>();
const onSubmit = vi.fn<() => Promise<undefined>>(() => pending.promise);
const input = submitDraft(onSubmit, "First draft");
await vi.waitFor(() => {
expect(onSubmit).toHaveBeenCalledOnce();
});
input.value = "Newer draft";
pending.reject(new Error("Send unavailable"));
await pending.promise.catch(() => undefined);

expect(input.value).toBe("Newer draft");
expect(input.dispatchEvent).not.toHaveBeenCalled();
});
});

function submitDraft(
onSubmit: ComponentProps<typeof PromptInput>["onSubmit"],
text: string
) {
// Exercise the actual rendered submit callback with a small DOM boundary.
// The native browser proof separately covers real FormData, textarea and events.
class Textarea {
value = text;
dispatchEvent = vi.fn<EventTarget["dispatchEvent"]>(() => true);
}
const input = new Textarea();
vi.stubGlobal("HTMLTextAreaElement", Textarea);
const NativeFormData = FormData;
vi.stubGlobal(
"FormData",
class extends NativeFormData {
constructor() {
super();
this.append("message", input.value);
}
}
);
const form = {
elements: { namedItem: () => input },
reset: () => {
input.value = "";
},
};
renderToStaticMarkup(
<PromptInput onSubmit={onSubmit}>
<PromptInputTextarea />
</PromptInput>
);
const { onSubmit: submit } = z
.object({
onSubmit: z.function({ input: [z.unknown()], output: z.void() }),
})
.parse(captured.forms.at(-1));
submit({
currentTarget: form,
preventDefault: vi.fn<() => void>(),
});
return input;
}