diff --git a/packages/ai/src/images-api-registry.ts b/packages/ai/src/images-api-registry.ts new file mode 100644 index 00000000..6db07dff --- /dev/null +++ b/packages/ai/src/images-api-registry.ts @@ -0,0 +1,76 @@ +import type { + AssistantImagesEventStream, + ImagesApi, + ImagesContext, + ImagesFunction, + ImagesModel, + ImagesOptions, +} from "./types.js"; + +export type ImagesApiFunction = ( + model: ImagesModel, + context: ImagesContext, + options?: ImagesOptions, +) => AssistantImagesEventStream; + +export interface ImagesApiProvider { + api: TApi; + images: ImagesFunction; +} + +interface ImagesApiProviderInternal { + api: ImagesApi; + images: ImagesApiFunction; +} + +type RegisteredImagesApiProvider = { + provider: ImagesApiProviderInternal; + sourceId?: string; +}; + +const imagesApiProviderRegistry = new Map(); + +function wrapImages( + api: TApi, + images: ImagesFunction, +): ImagesApiFunction { + return (model, context, options) => { + if (model.api !== api) { + throw new Error(`Mismatched api: ${model.api} expected ${api}`); + } + return images(model as ImagesModel, context, options as TOptions); + }; +} + +export function registerImagesApiProvider( + provider: ImagesApiProvider, + sourceId?: string, +): void { + imagesApiProviderRegistry.set(provider.api, { + provider: { + api: provider.api, + images: wrapImages(provider.api, provider.images), + }, + sourceId, + }); +} + +export function getImagesApiProvider(api: ImagesApi): ImagesApiProviderInternal | undefined { + return imagesApiProviderRegistry.get(api)?.provider; +} + +export function getImagesApiProviders(): ImagesApiProviderInternal[] { + return Array.from(imagesApiProviderRegistry.values(), (entry) => entry.provider); +} + +export function unregisterImagesApiProviders(sourceId: string): void { + for (const [api, entry] of imagesApiProviderRegistry.entries()) { + if (entry.sourceId === sourceId) { + imagesApiProviderRegistry.delete(api); + } + } +} + +export function clearImagesApiProviders(): void { + imagesApiProviderRegistry.clear(); +}