From 575c12503ce18fc4d353b491bed0e48b0211091c Mon Sep 17 00:00:00 2001 From: Jeremy Date: Thu, 1 May 2025 17:47:48 -0700 Subject: [PATCH] fix: support setting visibility on initial chat creation (#975) --- app/(chat)/api/chat/route.ts | 12 ++++-- app/(chat)/api/chat/schema.ts | 1 + app/(chat)/chat/[id]/page.tsx | 8 ++-- app/(chat)/page.tsx | 8 ++-- components/artifact.tsx | 6 +++ components/chat.tsx | 23 +++++++---- components/multimodal-input.tsx | 11 +++++- components/sidebar-history-item.tsx | 4 +- components/suggested-actions.tsx | 21 ++++++++-- components/visibility-selector.tsx | 7 ++-- hooks/use-chat-visibility.ts | 6 +-- lib/db/queries.ts | 4 ++ package.json | 2 +- tests/pages/chat.ts | 17 ++++++++ tests/routes/chat.test.ts | 60 ++++++++++++++++++++++++++++- 15 files changed, 158 insertions(+), 32 deletions(-) diff --git a/app/(chat)/api/chat/route.ts b/app/(chat)/api/chat/route.ts index 3b212ee..a6f6a48 100644 --- a/app/(chat)/api/chat/route.ts +++ b/app/(chat)/api/chat/route.ts @@ -49,7 +49,8 @@ export async function POST(request: Request) { } try { - const { id, message, selectedChatModel } = requestBody; + const { id, message, selectedChatModel, selectedVisibilityType } = + requestBody; const session = await auth(); @@ -80,7 +81,12 @@ export async function POST(request: Request) { message, }); - await saveChat({ id, userId: session.user.id, title }); + await saveChat({ + id, + userId: session.user.id, + title, + visibility: selectedVisibilityType, + }); } else { if (chat.userId !== session.user.id) { return new Response('Forbidden', { status: 403 }); @@ -236,7 +242,7 @@ export async function GET(request: Request) { return new Response('Not found', { status: 404 }); } - if (chat.userId !== session.user.id) { + if (chat.visibility === 'private' && chat.userId !== session.user.id) { return new Response('Forbidden', { status: 403 }); } diff --git a/app/(chat)/api/chat/schema.ts b/app/(chat)/api/chat/schema.ts index ae19df6..a452dc4 100644 --- a/app/(chat)/api/chat/schema.ts +++ b/app/(chat)/api/chat/schema.ts @@ -24,6 +24,7 @@ export const postRequestBodySchema = z.object({ .optional(), }), selectedChatModel: z.enum(['chat-model', 'chat-model-reasoning']), + selectedVisibilityType: z.enum(['public', 'private']), }); export type PostRequestBody = z.infer; diff --git a/app/(chat)/chat/[id]/page.tsx b/app/(chat)/chat/[id]/page.tsx index 9f3c3d8..c0024dc 100644 --- a/app/(chat)/chat/[id]/page.tsx +++ b/app/(chat)/chat/[id]/page.tsx @@ -60,8 +60,8 @@ export default async function Page(props: { params: Promise<{ id: string }> }) { }) { @@ -503,6 +507,8 @@ export const Artifact = memo(PureArtifact, (prevProps, nextProps) => { if (!equal(prevProps.votes, nextProps.votes)) return false; if (prevProps.input !== nextProps.input) return false; if (!equal(prevProps.messages, nextProps.messages.length)) return false; + if (prevProps.selectedVisibilityType !== nextProps.selectedVisibilityType) + return false; return true; }); diff --git a/components/chat.tsx b/components/chat.tsx index 9310882..b2c4e51 100644 --- a/components/chat.tsx +++ b/components/chat.tsx @@ -17,26 +17,32 @@ import { getChatHistoryPaginationKey } from './sidebar-history'; import { toast } from './toast'; import type { Session } from 'next-auth'; import { useSearchParams } from 'next/navigation'; +import { useChatVisibility } from '@/hooks/use-chat-visibility'; export function Chat({ id, initialMessages, - selectedChatModel, - selectedVisibilityType, + initialChatModel, + initialVisibilityType, isReadonly, session, autoResume, }: { id: string; initialMessages: Array; - selectedChatModel: string; - selectedVisibilityType: VisibilityType; + initialChatModel: string; + initialVisibilityType: VisibilityType; isReadonly: boolean; session: Session; autoResume: boolean; }) { const { mutate } = useSWRConfig(); + const { visibilityType } = useChatVisibility({ + chatId: id, + initialVisibilityType, + }); + const { messages, setMessages, @@ -57,7 +63,8 @@ export function Chat({ experimental_prepareRequestBody: (body) => ({ id, message: body.messages.at(-1), - selectedChatModel, + selectedChatModel: initialChatModel, + selectedVisibilityType: visibilityType, }), onFinish: () => { mutate(unstable_serialize(getChatHistoryPaginationKey)); @@ -109,8 +116,8 @@ export function Chat({
@@ -140,6 +147,7 @@ export function Chat({ messages={messages} setMessages={setMessages} append={append} + selectedVisibilityType={visibilityType} /> )} @@ -160,6 +168,7 @@ export function Chat({ reload={reload} votes={votes} isReadonly={isReadonly} + selectedVisibilityType={visibilityType} /> ); diff --git a/components/multimodal-input.tsx b/components/multimodal-input.tsx index b1f5b5c..f17372b 100644 --- a/components/multimodal-input.tsx +++ b/components/multimodal-input.tsx @@ -26,6 +26,7 @@ import type { UseChatHelpers } from '@ai-sdk/react'; import { AnimatePresence, motion } from 'framer-motion'; import { ArrowDown } from 'lucide-react'; import { useScrollToBottom } from '@/hooks/use-scroll-to-bottom'; +import type { VisibilityType } from './visibility-selector'; function PureMultimodalInput({ chatId, @@ -40,6 +41,7 @@ function PureMultimodalInput({ append, handleSubmit, className, + selectedVisibilityType, }: { chatId: string; input: UseChatHelpers['input']; @@ -53,6 +55,7 @@ function PureMultimodalInput({ append: UseChatHelpers['append']; handleSubmit: UseChatHelpers['handleSubmit']; className?: string; + selectedVisibilityType: VisibilityType; }) { const textareaRef = useRef(null); const { width } = useWindowSize(); @@ -220,7 +223,11 @@ function PureMultimodalInput({ {messages.length === 0 && attachments.length === 0 && uploadQueue.length === 0 && ( - + )} { const { visibilityType, setVisibilityType } = useChatVisibility({ chatId: chat.id, - initialVisibility: chat.visibility, + initialVisibilityType: chat.visibility, }); return ( diff --git a/components/suggested-actions.tsx b/components/suggested-actions.tsx index 535f2c2..430ab7a 100644 --- a/components/suggested-actions.tsx +++ b/components/suggested-actions.tsx @@ -3,14 +3,20 @@ import { motion } from 'framer-motion'; import { Button } from './ui/button'; import { memo } from 'react'; -import { UseChatHelpers } from '@ai-sdk/react'; +import type { UseChatHelpers } from '@ai-sdk/react'; +import type { VisibilityType } from './visibility-selector'; interface SuggestedActionsProps { chatId: string; append: UseChatHelpers['append']; + selectedVisibilityType: VisibilityType; } -function PureSuggestedActions({ chatId, append }: SuggestedActionsProps) { +function PureSuggestedActions({ + chatId, + append, + selectedVisibilityType, +}: SuggestedActionsProps) { const suggestedActions = [ { title: 'What are the advantages', @@ -71,4 +77,13 @@ function PureSuggestedActions({ chatId, append }: SuggestedActionsProps) { ); } -export const SuggestedActions = memo(PureSuggestedActions, () => true); +export const SuggestedActions = memo( + PureSuggestedActions, + (prevProps, nextProps) => { + if (prevProps.chatId !== nextProps.chatId) return false; + if (prevProps.selectedVisibilityType !== nextProps.selectedVisibilityType) + return false; + + return true; + }, +); diff --git a/components/visibility-selector.tsx b/components/visibility-selector.tsx index 7fd4059..b2d6dec 100644 --- a/components/visibility-selector.tsx +++ b/components/visibility-selector.tsx @@ -1,6 +1,6 @@ 'use client'; -import { ReactNode, useMemo, useState } from 'react'; +import { type ReactNode, useMemo, useState } from 'react'; import { Button } from '@/components/ui/button'; import { DropdownMenu, @@ -9,7 +9,6 @@ import { DropdownMenuTrigger, } from '@/components/ui/dropdown-menu'; import { cn } from '@/lib/utils'; - import { CheckCircleFillIcon, ChevronDownIcon, @@ -52,7 +51,7 @@ export function VisibilitySelector({ const { visibilityType, setVisibilityType } = useChatVisibility({ chatId, - initialVisibility: selectedVisibilityType, + initialVisibilityType: selectedVisibilityType, }); const selectedVisibility = useMemo( @@ -70,6 +69,7 @@ export function VisibilitySelector({ )} >