返回 presentation-ai
SlideGenerationContext.tsx
根目录 / src / components / notebook / presentation / editor / context / SlideGenerationContext.tsx
1 "use client";
2
3 import { useCompletion } from "@ai-sdk/react";
4 import {
5 createContext,
6 useCallback,
7 useContext,
8 useRef,
9 useState,
10 type ReactNode,
11 } from "react";
12
13 import { type ImageModelList } from "@/constants/image-models";
14 import { usePresentationState } from "@/states/presentation-state";
15 import { SlideParser } from "../../utils/parser";
16
17 function stripXmlCodeBlock(input: string): string {
18 let result = input.trim();
19 if (result.startsWith("```xml")) {
20 result = result.slice(6).trimStart();
21 }
22 if (result.endsWith("```")) {
23 result = result.slice(0, -3).trimEnd();
24 }
25 return result;
26 }
27
28 interface SlideGenerationContextValue {
29 isGenerating: boolean;
30 generatingSlideId: string | null;
31 generateSlide: (
32 slideId: string,
33 prompt: string,
34 options?: SlideGenerationOptions,
35 ) => void;
36 cancelGeneration: () => void;
37 }
38
39 type SlideGenerationOptions = {
40 slideType?: "standard" | "image";
41 imageStyle?: string;
42 textDensity?: string;
43 imageModel?: ImageModelList;
44 };
45
46 const SlideGenerationContext =
47 createContext<SlideGenerationContextValue | null>(null);
48
49 export function SlideGenerationProvider({ children }: { children: ReactNode }) {
50 const [generatingSlideId, setGeneratingSlideId] = useState<string | null>(
51 null,
52 );
53
54 // Refs for parsing
55 const parserRef = useRef<SlideParser | null>(null);
56 if (!parserRef.current) {
57 parserRef.current = new SlideParser();
58 }
59 const parser = parserRef.current;
60 const slideIdRef = useRef<string | null>(null);
61 const slideTypeRef = useRef<SlideGenerationOptions["slideType"]>("standard");
62 const imageModelRef = useRef<SlideGenerationOptions["imageModel"]>(undefined);
63
64 const language = usePresentationState((s) => s.language);
65 const setSlides = usePresentationState((s) => s.setSlides);
66 const startRootImageGeneration = usePresentationState(
67 (s) => s.startRootImageGeneration,
68 );
69
70 const { complete, isLoading, stop } = useCompletion({
71 api: "/api/presentation/generate-slide",
72 onFinish: (_prompt, finalCompletion) => {
73 // Parse final content and update the slide only when fully complete
74 const processedCompletion = stripXmlCodeBlock(finalCompletion);
75 parser.reset();
76 parser.parseChunk(processedCompletion);
77 parser.finalize();
78 const parsedSlides = parser.getAllSlides();
79
80 if (parsedSlides.length > 0 && slideIdRef.current) {
81 const generatedSlide = parsedSlides[0];
82 const targetSlideId = slideIdRef.current;
83 const isImageSlide = slideTypeRef.current === "image";
84
85 // Update the target slide with generated content
86 const currentSlides = usePresentationState.getState().slides;
87 const updatedSlides = currentSlides.map((slide) => {
88 if (slide.id === targetSlideId) {
89 return {
90 ...slide,
91 content: generatedSlide!.content,
92 layoutType: generatedSlide!.layoutType ?? slide.layoutType,
93 rootImage: generatedSlide!.rootImage ?? slide.rootImage,
94 isImageSlide,
95 };
96 }
97 return slide;
98 });
99
100 setSlides(updatedSlides);
101
102 if (generatedSlide?.rootImage?.query) {
103 startRootImageGeneration(
104 targetSlideId,
105 generatedSlide.rootImage.query,
106 {
107 imageModel: imageModelRef.current,
108 source: isImageSlide ? "ai" : undefined,
109 },
110 );
111 }
112
113 }
114
115 // Reset state
116 setGeneratingSlideId(null);
117 slideIdRef.current = null;
118 slideTypeRef.current = "standard";
119 imageModelRef.current = undefined;
120 parser.reset();
121 },
122 onError: (error) => {
123 console.error("Failed to generate slide:", error);
124 setGeneratingSlideId(null);
125 slideIdRef.current = null;
126 parser.reset();
127 },
128 });
129
130 const generateSlide = useCallback(
131 (slideId: string, prompt: string, options?: SlideGenerationOptions) => {
132 if (isLoading) return;
133
134 // Reset state
135 setGeneratingSlideId(slideId);
136 parser.reset();
137 slideIdRef.current = slideId;
138 slideTypeRef.current = options?.slideType ?? "standard";
139 imageModelRef.current = options?.imageModel;
140
141 // Get current slide for context
142 const slides = usePresentationState.getState().slides;
143 const currentSlide = slides.find((s) => s.id === slideId);
144
145 // Start generation
146 void complete(prompt, {
147 body: {
148 prompt,
149 language,
150 currentSlide: currentSlide
151 ? JSON.stringify(currentSlide.content)
152 : undefined,
153 slideType: options?.slideType,
154 imageStyle: options?.imageStyle,
155 textDensity: options?.textDensity,
156 imageModel: options?.imageModel,
157 },
158 });
159 },
160 [isLoading, complete, language],
161 );
162
163 const cancelGeneration = useCallback(() => {
164 stop();
165 setGeneratingSlideId(null);
166 slideIdRef.current = null;
167 slideTypeRef.current = "standard";
168 imageModelRef.current = undefined;
169 parser.reset();
170 }, [stop]);
171
172 return (
173 <SlideGenerationContext.Provider
174 value={{
175 isGenerating: isLoading,
176 generatingSlideId,
177 generateSlide,
178 cancelGeneration,
179 }}
180 >
181 {children}
182 </SlideGenerationContext.Provider>
183 );
184 }
185
186 export function useSlideGeneration(): SlideGenerationContextValue {
187 const context = useContext(SlideGenerationContext);
188 if (!context) {
189 throw new Error(
190 "useSlideGeneration must be used within a SlideGenerationProvider",
191 );
192 }
193 return context;
194 }
195
195 lines Plain Text