返回 presentation-ai
image-generation.ts
根目录 / src / lib / presentation / image-generation.ts
1 import {
2 type PlateNode,
3 type PlateSlide,
4 type RootImage,
5 } from "@/components/notebook/presentation/utils/parser";
6 import { type ImageModelList } from "@/constants/image-models";
7
8 type PresentationImageStockProvider = "unsplash" | "pixabay" | "google";
9
10 export type PresentationImageGenerationSource = "ai" | "stock" | "gif";
11
12 export type PresentationImageGenerationTarget =
13 | {
14 kind: "root";
15 slideId: string;
16 }
17 | {
18 elementId: string;
19 kind: "element";
20 slideId: string;
21 };
22
23 type PresentationImageGenerationStatus =
24 | "queued"
25 | "generating"
26 | "success"
27 | "error";
28
29 export interface PresentationImageGenerationJob {
30 error?: string;
31 imageModel?: ImageModelList;
32 presentationId?: string;
33 query: string;
34 source: PresentationImageGenerationSource;
35 status: PresentationImageGenerationStatus;
36 stockImageProvider?: PresentationImageStockProvider;
37 target: PresentationImageGenerationTarget;
38 url?: string;
39 }
40
41 export interface PresentationImageTargetState {
42 elementId?: string;
43 imageSource?: RootImage["imageSource"];
44 isImageSlide?: boolean;
45 layoutType?: RootImage["layoutType"];
46 query?: string;
47 stockImageProvider?: PresentationImageStockProvider;
48 url?: string;
49 }
50
51 export interface PresentationPendingImageTarget {
52 state: PresentationImageTargetState;
53 target: PresentationImageGenerationTarget;
54 }
55
56 export function getPresentationImageGenerationKey(
57 target: PresentationImageGenerationTarget,
58 ): string {
59 return target.kind === "root"
60 ? target.slideId
61 : `${target.slideId}:${target.elementId}`;
62 }
63
64 export function getRootImageGenerationTarget(
65 slideId: string,
66 ): PresentationImageGenerationTarget {
67 return { kind: "root", slideId };
68 }
69
70 export function getElementImageGenerationTarget(
71 slideId: string,
72 elementId: string,
73 ): PresentationImageGenerationTarget {
74 return { kind: "element", slideId, elementId };
75 }
76
77 export function getElementImageGenerationKey(
78 slideId: string,
79 elementId: string,
80 ): string {
81 return getPresentationImageGenerationKey(
82 getElementImageGenerationTarget(slideId, elementId),
83 );
84 }
85
86 export function resolvePresentationImageGenerationSource({
87 globalImageSource,
88 imageSource,
89 isImageSlide,
90 }: {
91 globalImageSource: "automatic" | "ai" | "stock" | "gif";
92 imageSource?: RootImage["imageSource"];
93 isImageSlide?: boolean;
94 }): PresentationImageGenerationSource {
95 if (imageSource === "gif") {
96 return "gif";
97 }
98
99 if (imageSource === "search") {
100 return "stock";
101 }
102
103 if (imageSource === "generate") {
104 return "ai";
105 }
106
107 if (isImageSlide) {
108 return "ai";
109 }
110
111 if (globalImageSource === "gif") {
112 return "gif";
113 }
114
115 if (globalImageSource === "ai") {
116 return "ai";
117 }
118
119 return "stock";
120 }
121
122 export function getGeneratedPresentationImageSource(
123 source: PresentationImageGenerationSource,
124 ): RootImage["imageSource"] {
125 if (source === "gif") {
126 return "gif";
127 }
128
129 return source === "stock" ? "search" : "generate";
130 }
131
132 function isRecord(value: unknown): value is Record<string, unknown> {
133 return typeof value === "object" && value !== null && !Array.isArray(value);
134 }
135
136 function findImageElementState(
137 nodes: PlateNode[],
138 elementId: string,
139 ): PresentationImageTargetState | undefined {
140 for (const node of nodes) {
141 if (!isRecord(node)) {
142 continue;
143 }
144
145 if (node.id === elementId) {
146 return {
147 imageSource:
148 node.imageSource === "generate" ||
149 node.imageSource === "search" ||
150 node.imageSource === "gif" ||
151 node.imageSource === "upload"
152 ? node.imageSource
153 : undefined,
154 query:
155 typeof node.query === "string"
156 ? node.query
157 : typeof node.prompt === "string"
158 ? node.prompt
159 : undefined,
160 stockImageProvider:
161 node.stockImageProvider === "unsplash" ||
162 node.stockImageProvider === "pixabay" ||
163 node.stockImageProvider === "google"
164 ? node.stockImageProvider
165 : undefined,
166 url: typeof node.url === "string" ? node.url : undefined,
167 };
168 }
169
170 if (Array.isArray(node.children)) {
171 const childResult = findImageElementState(
172 node.children as PlateNode[],
173 elementId,
174 );
175
176 if (childResult) {
177 return childResult;
178 }
179 }
180 }
181
182 return undefined;
183 }
184
185 function collectImageElementStates(
186 nodes: PlateNode[],
187 slideId: string,
188 results: PresentationPendingImageTarget[],
189 ): void {
190 for (const node of nodes) {
191 if (!isRecord(node)) {
192 continue;
193 }
194
195 const imageQuery =
196 node.type === "icon-item"
197 ? typeof node.prompt === "string"
198 ? node.prompt
199 : undefined
200 : typeof node.query === "string"
201 ? node.query
202 : typeof node.prompt === "string"
203 ? node.prompt
204 : undefined;
205 const isImageGenerationElement =
206 node.type === "img" || node.type === "image" || node.type === "icon-item";
207
208 if (
209 typeof node.id === "string" &&
210 isImageGenerationElement &&
211 imageQuery !== undefined &&
212 imageQuery.trim().length > 0 &&
213 typeof node.url !== "string"
214 ) {
215 results.push({
216 state: {
217 elementId: node.id,
218 imageSource:
219 node.imageSource === "generate" ||
220 node.imageSource === "search" ||
221 node.imageSource === "gif" ||
222 node.imageSource === "upload"
223 ? node.imageSource
224 : undefined,
225 query: imageQuery,
226 stockImageProvider:
227 node.stockImageProvider === "unsplash" ||
228 node.stockImageProvider === "pixabay" ||
229 node.stockImageProvider === "google"
230 ? node.stockImageProvider
231 : undefined,
232 url: undefined,
233 },
234 target: getElementImageGenerationTarget(slideId, node.id),
235 });
236 }
237
238 if (Array.isArray(node.children)) {
239 collectImageElementStates(node.children as PlateNode[], slideId, results);
240 }
241 }
242 }
243
244 export function getPendingPresentationImageTargets(
245 slides: PlateSlide[],
246 ): PresentationPendingImageTarget[] {
247 const results: PresentationPendingImageTarget[] = [];
248
249 for (const slide of slides) {
250 if (
251 slide.rootImage?.query &&
252 !slide.rootImage.url &&
253 !slide.rootImage.isQueryStreaming
254 ) {
255 results.push({
256 state: {
257 imageSource: slide.rootImage.imageSource,
258 isImageSlide: slide.isImageSlide,
259 layoutType: slide.rootImage.layoutType ?? slide.layoutType,
260 query: slide.rootImage.query,
261 stockImageProvider: slide.rootImage.stockImageProvider,
262 },
263 target: getRootImageGenerationTarget(slide.id),
264 });
265 }
266
267 collectImageElementStates(slide.content, slide.id, results);
268 }
269
270 return results;
271 }
272
273 export function getPresentationImageTargetState(
274 slides: PlateSlide[],
275 target: PresentationImageGenerationTarget,
276 ): PresentationImageTargetState | undefined {
277 const slide = slides.find((candidate) => candidate.id === target.slideId);
278 if (!slide) {
279 return undefined;
280 }
281
282 if (target.kind === "root") {
283 return slide.rootImage
284 ? {
285 imageSource: slide.rootImage.imageSource,
286 isImageSlide: slide.isImageSlide,
287 layoutType: slide.rootImage.layoutType ?? slide.layoutType,
288 query: slide.rootImage.query,
289 stockImageProvider: slide.rootImage.stockImageProvider,
290 url: slide.rootImage.url,
291 }
292 : undefined;
293 }
294
295 return findImageElementState(slide.content, target.elementId);
296 }
297
298 function updateImageElementNodes(
299 nodes: PlateNode[],
300 elementId: string,
301 patch: Record<string, unknown>,
302 ): {
303 changed: boolean;
304 nodes: PlateNode[];
305 } {
306 let changed = false;
307
308 const nextNodes = nodes.map((node) => {
309 if (!isRecord(node)) {
310 return node;
311 }
312
313 let nextNode: Record<string, unknown> = node;
314
315 if (node.id === elementId) {
316 changed = true;
317 const promptPatch =
318 node.type === "icon-item" && typeof patch.query === "string"
319 ? { prompt: patch.query }
320 : {};
321 nextNode = {
322 ...node,
323 ...patch,
324 ...promptPatch,
325 };
326 } else if (Array.isArray(node.children)) {
327 const childResult = updateImageElementNodes(
328 node.children as PlateNode[],
329 elementId,
330 patch,
331 );
332
333 if (childResult.changed) {
334 changed = true;
335 nextNode = {
336 ...node,
337 children: childResult.nodes,
338 };
339 }
340 }
341
342 return nextNode as PlateNode;
343 });
344
345 return {
346 changed,
347 nodes: changed ? nextNodes : nodes,
348 };
349 }
350
351 export function updatePresentationImageTarget(
352 slides: PlateSlide[],
353 target: PresentationImageGenerationTarget,
354 patch: Partial<RootImage> & Record<string, unknown>,
355 ): PlateSlide[] {
356 return slides.map((slide) => {
357 if (slide.id !== target.slideId) {
358 return slide;
359 }
360
361 if (target.kind === "root") {
362 return {
363 ...slide,
364 rootImage: {
365 ...(slide.rootImage ?? { query: "" }),
366 ...patch,
367 },
368 };
369 }
370
371 const result = updateImageElementNodes(slide.content, target.elementId, {
372 ...patch,
373 });
374
375 return result.changed
376 ? {
377 ...slide,
378 content: result.nodes,
379 }
380 : slide;
381 });
382 }
383
383 lines TYPESCRIPT