diff --git a/packages/ui/src/components/Command/AiCommand.tsx b/packages/ui/src/components/Command/AiCommand.tsx
index c4245890e68..95e94e3f22d 100644
--- a/packages/ui/src/components/Command/AiCommand.tsx
+++ b/packages/ui/src/components/Command/AiCommand.tsx
@@ -18,7 +18,7 @@ import { SSE } from 'sse.js'
import { Alert, Button, IconAlertTriangle, IconCornerDownLeft, IconUser, Input } from 'ui'
import { AiIcon, AiIconChat } from './Command.icons'
-import { CommandGroup, CommandItem } from './Command.utils'
+import { CommandGroup, CommandItem, useAutoInputFocus, useHistoryKeys } from './Command.utils'
import { useCommandMenu } from './CommandMenuProvider'
import { AiWarning } from './Command.alerts'
@@ -345,6 +345,16 @@ const AiCommand = () => {
setIsLoading,
})
+ const inputRef = useAutoInputFocus()
+
+ useHistoryKeys({
+ enable: !isResponding,
+ messages: messages
+ .filter(({ role }) => role === MessageRole.User)
+ .map(({ content }) => content),
+ setPrompt: setSearch,
+ })
+
const handleSubmit = useCallback(
(message: string) => {
setSearch('')
@@ -459,6 +469,7 @@ const AiCommand = () => {
{messages.length > 0 && !hasError && }
void
+}
+
+/**
+ * Enables a shell-style message history when hitting
+ * up/down on the keyboard
+ */
+export function useHistoryKeys({ enable, messages, setPrompt }: UseHistoryKeysOptions) {
+ // Message index when hitting up/down on the keyboard (shell style)
+ const [, setMessageSelectionIndex] = React.useState(0)
+
+ React.useEffect(() => {
+ if (enable) {
+ return
+ }
+
+ // Note: intentionally setting index to 1 greater than max index
+ setMessageSelectionIndex(messages.length)
+ }, [messages, enable])
+
+ React.useEffect(() => {
+ function onKeyDown(e: KeyboardEvent) {
+ switch (e.key) {
+ case 'ArrowUp':
+ setMessageSelectionIndex((index) => {
+ const newIndex = Math.max(index - 1, 0)
+ setPrompt(messages[newIndex] ?? '')
+ return newIndex
+ })
+ return
+ case 'ArrowDown':
+ setMessageSelectionIndex((index) => {
+ const newIndex = Math.min(index + 1, messages.length)
+ setPrompt(messages[newIndex] ?? '')
+ return newIndex
+ })
+ return
+ default:
+ return
+ }
+ }
+
+ window.addEventListener('keydown', onKeyDown)
+
+ return () => {
+ window.removeEventListener('keydown', onKeyDown)
+ }
+ }, [messages])
+}
+
+/**
+ * Automatically focuses an input on key press
+ * and on load (after the call stack)
+ *
+ * @returns An input ref for the input to focus
+ */
+export function useAutoInputFocus() {
+ const [input, setInput] = React.useState()
+
+ // Use a callback-style ref to access the element when it mounts
+ const inputRef = React.useCallback((inputElement: HTMLInputElement) => {
+ if (inputElement) {
+ setInput(inputElement)
+
+ // We need to delay the focus until the end of the call stack
+ // due to order of operations
+ setTimeout(() => {
+ inputElement.focus()
+ }, 0)
+ }
+ }, [])
+
+ // Focus the input when typing from anywhere
+ React.useEffect(() => {
+ function onKeyDown() {
+ input?.focus()
+ }
+
+ window.addEventListener('keydown', onKeyDown)
+
+ return () => {
+ window.removeEventListener('keydown', onKeyDown)
+ }
+ }, [input])
+
+ return inputRef
+}
diff --git a/packages/ui/src/components/Command/GenerateSQL/GenerateSQL.tsx b/packages/ui/src/components/Command/GenerateSQL/GenerateSQL.tsx
index 2f7f81915b1..e6911f05822 100644
--- a/packages/ui/src/components/Command/GenerateSQL/GenerateSQL.tsx
+++ b/packages/ui/src/components/Command/GenerateSQL/GenerateSQL.tsx
@@ -18,7 +18,7 @@ import {
import { cn } from '../../../utils/cn'
import { AiIcon, AiIconChat } from '../Command.icons'
-import { CommandItem } from '../Command.utils'
+import { CommandItem, useAutoInputFocus, useHistoryKeys } from '../Command.utils'
import { useCommandMenu } from '../CommandMenuProvider'
import { SAMPLE_QUERIES } from '../Command.constants'
import SQLOutputActions from './SQLOutputActions'
@@ -47,6 +47,16 @@ const GenerateSQL = () => {
setIsLoading,
})
+ const inputRef = useAutoInputFocus()
+
+ useHistoryKeys({
+ enable: !isResponding,
+ messages: messages
+ .filter(({ role }) => role === MessageRole.User)
+ .map(({ content }) => content),
+ setPrompt: setSearch,
+ })
+
const handleSubmit = useCallback(
(message: string) => {
setSearch('')
@@ -271,15 +281,7 @@ const GenerateSQL = () => {
)}
{
- if (inputElement) {
- // We need to delay the focus until the end of the call stack
- // due to order of operations
- setTimeout(() => {
- inputElement.focus()
- }, 0)
- }
- }}
+ inputRef={inputRef}
className="bg-scale-100 rounded mx-3 mb-4"
autoFocus
placeholder={