import { Image, ImageGenerateParams } from "openai/resources"; import { VisionInterface } from "../../interfaces/Vision"; import { ImageParams, ImageResult, VideoParams, VideoResult } from "../../types/vision/vision"; import { ErnieAPI } from "../../utils/erine"; import { BaiduTextToImageBasicRequest, BaiduTextToImageSpeechRequest } from "../../types/vision/baidu"; import { Logger } from "../../utils/Logger"; import { VisionResultStream } from "../../utils/stream"; export class VisionBaiduCloudService implements VisionInterface { private clientId: string; private clientSecret: string; constructor(clientId: string, clientSecret: string) { this.clientId = clientId || ""; this.clientSecret = clientSecret || ""; if (!this.clientId || !this.clientSecret) { throw new Error("clientId and clientSecret is required"); } } async image(params: ImageGenerateParams & ImageParams): Promise> { if (params.model === "speed") { return await this.imageSpeed(params) } else if (params.model === "basic") { return await this.imageBasic(params) } else if (params.model === "pro") { return await this.imageSpeed(params) } else { return Promise.reject(new Error("model is required")) } } async imageSpeed(params: ImageGenerateParams & ImageParams): Promise> { const api = await ErnieAPI.getInstance(this.clientId, this.clientSecret) const stream = new VisionResultStream() let url = "/rpc/2.0/wenxin/v1/extreme/textToImage" if (params.model !== "speed") { url = "/rpc/2.0/ernievilg/v1/txt2imgv2" } const req: BaiduTextToImageSpeechRequest = { prompt: params.prompt, width: params.width || 1024, height: params.height || 1024, image_num: params.n || 1, } if (params.image_url) { if (params.image_url.startsWith("https") || params.image_url.startsWith("http")) { req.url = params.image_url } else { req.image = params.image_url } req.change_degree = params.change_degree || 5 } const res = await api.request({ url, method: "POST", data: req }) if (!res.ok || res.status !== 200) { this.streamFail(stream, res.statusText) return stream } const json = await res.json() const taskId = json.data.task_id if (!taskId) { this.streamFail(stream, "task_id is empty") return stream } let taskUrl = "/rpc/2.0/wenxin/v1/extreme/getImg" if (params.model !== "speed") { taskUrl = "/rpc/2.0/ernievilg/v1/getImgv2" } const taskRequest = async (taskId: string): Promise => { const taskRes = await api.request({ url: taskUrl, method: "POST", data: { task_id: taskId } }) if (!taskRes.ok || taskRes.status !== 200) { this.streamFail(stream, taskRes.statusText) return } const taskJson = await taskRes.json() const taskData = taskJson.data Logger.debug("taskData:", JSON.stringify(taskData, null, 2)) if (taskData.task_progress === 1) { stream.write({ status: "completed", created_at: Date.now(), data: taskData.sub_task_result_list.map((item: any) => { return { format: "url", url: item.final_image_list[0].img_url } }) } as ImageResult) stream.end() return } else { stream.write({ status: "processing", created_at: Date.now() } as ImageResult) setTimeout(() => { taskRequest(taskId) }, 1000) } } taskRequest(taskId) return stream } private streamFail(stream: VisionResultStream, message: string) { stream.write({ status: "failed", created_at: Date.now(), message } as ImageResult) stream.end() } async imageBasic(params: ImageGenerateParams & ImageParams): Promise> { const api = await ErnieAPI.getInstance(this.clientId, this.clientSecret) const stream = new VisionResultStream() const url = "/rpc/2.0/wenxin/v1/basic/textToImage" const req: BaiduTextToImageBasicRequest = { text: params.prompt, resolution: `${params.width}*${params.height}`, num: params.n || 1 } const res = await api.request({ url, method: "POST", data: req }) if (!res.ok || res.status !== 200) { this.streamFail(stream, res.statusText) return stream } const json = await res.json() const taskId = json.data.taskId if (!taskId) { this.streamFail(stream, "task_id is empty") return stream } const taskUrl = "/rpc/2.0/wenxin/v1/basic/getImg" const taskRequest = async (taskId: string): Promise => { const taskRes = await api.request({ url: taskUrl, method: "POST", data: { taskId: taskId } }) if (!taskRes.ok || taskRes.status !== 200) { this.streamFail(stream, taskRes.statusText) return } const taskJson = await taskRes.json() const taskData = taskJson.data Logger.debug("taskData:", taskData) if (taskData.status === 1) { stream.write({ status: "completed", created_at: Date.now(), data: taskData.imgUrls.map((item: any) => { return { format: "url", url: item.image } }) } as ImageResult) stream.end() return } else { stream.write({ status: "processing", created_at: Date.now(), message: "is waiting : " + taskData.waiting } as ImageResult) setTimeout(() => { taskRequest(taskId) }, 1000) } } taskRequest(taskId) return stream } // async imagePro(params: ImageGenerateParams & ImageParams): Promise> { // const url = "/rpc/2.0/ernievilg/v1/txt2imgv2" // } video(params: VideoParams): Promise> { throw new Error("Method not implemented."); } }