diff --git a/open-sse/executors/kie.ts b/open-sse/executors/kie.ts index fbb2758822..494c867e60 100644 --- a/open-sse/executors/kie.ts +++ b/open-sse/executors/kie.ts @@ -1,5 +1,13 @@ import { BaseExecutor } from "./base.ts"; import { sleep } from "../utils/sleep.ts"; +import { + isJsonObject, + normalizeKieTaskState, + type JsonObject, + type KieTaskState, +} from "../utils/kieTask.ts"; + +export type { KieTaskState } from "../utils/kieTask.ts"; type KieTaskInput = { baseUrl: string; @@ -16,10 +24,8 @@ type KiePollInput = { pollIntervalMs: number; }; -export type KieTaskState = "success" | "failed" | "pending"; - export type KieTaskRecord = { - data: any; + data: JsonObject; state: KieTaskState; }; @@ -27,45 +33,6 @@ function normalizeBaseUrl(baseUrl: string): string { return baseUrl.replace(/\/$/, ""); } -export function normalizeKieTaskState(recordData: any): KieTaskState { - const state = String( - recordData?.data?.status ?? - recordData?.data?.state ?? - recordData?.data?.successFlag ?? - recordData?.msg ?? - "PENDING" - ).toUpperCase(); - - if ( - state === "SUCCESS" || - state === "1" || - state === "FINISHED" || - state === "COMPLETE" || - state === "COMPLETED" || - state === "FIRST_SUCCESS" || - state === "ALL_SUCCESS" || - state.includes("SUCCESS") - ) { - return "success"; - } - - if ( - state === "FAIL" || - state === "FAILED" || - state === "ERROR" || - state === "2" || - state === "3" || - state.includes("FAIL") || - state.includes("ERROR") || - state === "CREATE_TASK_FAILED" || - state === "GENERATE_FAILED" - ) { - return "failed"; - } - - return "pending"; -} - export class KieExecutor extends BaseExecutor { constructor() { super("kie", { baseUrl: "https://api.kie.ai" }); @@ -79,7 +46,7 @@ export class KieExecutor extends BaseExecutor { return `${normalizeBaseUrl(baseUrl)}/api/v1/jobs/recordInfo`; } - async createTask({ baseUrl, token, payload, endpoint }: KieTaskInput): Promise { + async createTask({ baseUrl, token, payload, endpoint }: KieTaskInput): Promise { const res = await fetch(this.getTaskCreateUrl(baseUrl, endpoint), { method: "POST", headers: { @@ -96,7 +63,8 @@ export class KieExecutor extends BaseExecutor { }); } - return res.json(); + const data = (await res.json()) as unknown; + return isJsonObject(data) ? data : {}; } async pollTask({ @@ -124,10 +92,11 @@ export class KieExecutor extends BaseExecutor { }); } - const data = await res.json(); - const state = normalizeKieTaskState(data); + const data = (await res.json()) as unknown; + const recordData = isJsonObject(data) ? data : {}; + const state = normalizeKieTaskState(recordData); if (state !== "pending") { - return { data, state }; + return { data: recordData, state }; } await sleep(pollIntervalMs); diff --git a/open-sse/handlers/audioSpeech.ts b/open-sse/handlers/audioSpeech.ts index 27d0853d6c..be07120491 100644 --- a/open-sse/handlers/audioSpeech.ts +++ b/open-sse/handlers/audioSpeech.ts @@ -21,6 +21,13 @@ import { getSpeechProvider, parseSpeechModel } from "../config/audioRegistry.ts" import { buildAuthHeaders } from "../config/registryUtils.ts"; import { kieExecutor } from "../executors/kie.ts"; import { errorResponse } from "../utils/error.ts"; +import { + getKieCallbackUrl, + getKieErrorMessage, + getKieErrorStatus, + isJsonObject, + parseKieResultJson, +} from "../utils/kieTask.ts"; /** * Return a CORS error response from an upstream fetch failure @@ -69,15 +76,6 @@ function audioStreamResponse(res, defaultContentType = "audio/mpeg") { }); } -function getKieCallbackUrl(body: any): string { - return ( - body.callBackUrl || - body.callback_url || - body.callbackUrl || - "https://omniroute.local/api/kie/callback" - ); -} - function normalizeKieElevenLabsVoice(voice: unknown): string { const value = typeof voice === "string" ? voice.trim() : ""; const aliases: Record = { @@ -91,17 +89,7 @@ function normalizeKieElevenLabsVoice(voice: unknown): string { return aliases[value.toLowerCase()] || value || "Rachel"; } -function parseKieResultJson(recordData: any): any { - try { - return typeof recordData?.data?.resultJson === "string" - ? JSON.parse(recordData.data.resultJson) - : recordData?.data?.resultJson || {}; - } catch { - return {}; - } -} - -function findAudioUrlDeep(value: any): string | null { +function findAudioUrlDeep(value: unknown): string | null { if (!value) return null; if (typeof value === "string") { @@ -119,7 +107,7 @@ function findAudioUrlDeep(value: any): string | null { return null; } - if (typeof value === "object") { + if (isJsonObject(value)) { const preferredKeys = [ "audio_url", "audioUrl", @@ -145,16 +133,20 @@ function findAudioUrlDeep(value: any): string | null { return null; } -function findKieAudioUrl(recordData: any): string | null { +function findKieAudioUrl(recordData: unknown): string | null { + const record = isJsonObject(recordData) ? recordData : {}; + const data = isJsonObject(record.data) ? record.data : {}; const resultJson = parseKieResultJson(recordData); + const response = data.response; + const nestedData = data.data; const candidates = [ - recordData?.data?.response, - recordData?.data, + response, + data, resultJson, - ...(Array.isArray(recordData?.data?.response) ? recordData.data.response : []), - ...(Array.isArray(recordData?.data?.data) ? recordData.data.data : []), - ...(Array.isArray(resultJson?.data) ? resultJson.data : []), - ...(Array.isArray(resultJson?.result) ? resultJson.result : []), + ...(Array.isArray(response) ? response : []), + ...(Array.isArray(nestedData) ? nestedData : []), + ...(Array.isArray(resultJson.data) ? resultJson.data : []), + ...(Array.isArray(resultJson.result) ? resultJson.result : []), ]; for (const item of candidates) { @@ -453,13 +445,14 @@ async function handleKieAudioSpeech(providerConfig, body, modelId, token) { token, payload, }); - } catch (err: any) { + } catch (err: unknown) { + const status = getKieErrorStatus(err, 502); return Response.json( { - error: { message: err?.message || "Kie audio createTask failed", code: err?.status || 502 }, + error: { message: getKieErrorMessage(err, "Kie audio createTask failed"), code: status }, }, { - status: Number(err?.status) || 502, + status, headers: { "Access-Control-Allow-Origin": getCorsOrigin() }, } ); @@ -504,10 +497,10 @@ async function pollKieAudioResult(baseUrl, modelId, taskId, token) { } return errorResponse(502, "Kie audio task completed without audio URL"); } - } catch (err: any) { + } catch (err: unknown) { return errorResponse( - Number(err?.status) || 504, - err?.message || "Kie audio generation timed out or failed" + getKieErrorStatus(err, 504), + getKieErrorMessage(err, "Kie audio generation timed out or failed") ); } diff --git a/open-sse/handlers/audioTranscription.ts b/open-sse/handlers/audioTranscription.ts index 9a6a9b2e87..d2e512c83f 100644 --- a/open-sse/handlers/audioTranscription.ts +++ b/open-sse/handlers/audioTranscription.ts @@ -306,16 +306,20 @@ async function handleKieAudioTranscription(providerConfig, file, modelId, token) }, }, }); - } catch (err: any) { + } catch (err: unknown) { + const status = + typeof err === "object" && err !== null && "status" in err + ? Number((err as { status?: unknown }).status) || 502 + : 502; return Response.json( { error: { - message: err?.message || "Kie transcription createTask failed", - code: err?.status || 502, + message: err instanceof Error ? err.message : "Kie transcription createTask failed", + code: status, }, }, { - status: Number(err?.status) || 502, + status, headers: { "Access-Control-Allow-Origin": getCorsOrigin() }, } ); @@ -359,10 +363,14 @@ async function pollKieTranscriptionResult(baseUrl, modelId, taskId, token) { { headers: { "Access-Control-Allow-Origin": getCorsOrigin() } } ); } - } catch (err: any) { + } catch (err: unknown) { + const status = + typeof err === "object" && err !== null && "status" in err + ? Number((err as { status?: unknown }).status) || 504 + : 504; return errorResponse( - Number(err?.status) || 504, - err?.message || "Kie transcription generation timed out or failed" + status, + err instanceof Error ? err.message : "Kie transcription generation timed out or failed" ); } diff --git a/open-sse/handlers/imageGeneration.ts b/open-sse/handlers/imageGeneration.ts index 03a8e45ecd..63d7bd5c55 100644 --- a/open-sse/handlers/imageGeneration.ts +++ b/open-sse/handlers/imageGeneration.ts @@ -21,6 +21,12 @@ import { kieExecutor } from "../executors/kie.ts"; import { mapImageSize } from "../translator/image/sizeMapper.ts"; import { saveCallLog } from "@/lib/usageDb"; import { sleep } from "../utils/sleep.ts"; +import { + getKieErrorMessage, + getKieErrorStatus, + isJsonObject, + parseKieResultJson, +} from "../utils/kieTask.ts"; import { submitComfyWorkflow, pollComfyResult, @@ -31,10 +37,25 @@ import { interface KieImageOptions { model: string; provider: string; - providerConfig: any; - body: any; - credentials: any; - log: any; + providerConfig: { + baseUrl: string; + statusUrl?: string; + }; + body: Record & { + prompt?: unknown; + size?: unknown; + n?: unknown; + timeout_ms?: unknown; + poll_interval_ms?: unknown; + }; + credentials?: { + apiKey?: string; + accessToken?: string; + } | null; + log?: { + info: (scope: string, message: string) => void; + error: (scope: string, message: string) => void; + } | null; } const OPENAI_IMAGE_TO_IMAGE_MODELS = new Set([ @@ -300,20 +321,14 @@ export async function handleImageGeneration({ body, credentials, log, resolvedPr return handleOpenAIImageGeneration({ model, provider, providerConfig, body, credentials, log }); } -function normalizeKieImageResult(recordData: any): string[] { - let resultJson: Record = {}; - try { - resultJson = - typeof recordData?.data?.resultJson === "string" - ? JSON.parse(recordData.data.resultJson) - : recordData?.data?.resultJson || {}; - } catch { - resultJson = {}; - } - +function normalizeKieImageResult(recordData: unknown): string[] { + const record = isJsonObject(recordData) ? recordData : {}; + const data = isJsonObject(record.data) ? record.data : {}; + const response = isJsonObject(data.response) ? data.response : {}; + const resultJson = parseKieResultJson(recordData); const urls = new Set(); - const add = (val: any) => { + const add = (val: unknown) => { if (typeof val === "string" && val.startsWith("http")) urls.add(val); if (Array.isArray(val)) { val.forEach((v) => { @@ -329,13 +344,13 @@ function normalizeKieImageResult(recordData: any): string[] { add(resultJson?.imageUrl); // Check data.response (common in 4o-image API) - add(recordData?.data?.response?.resultUrls); - add(recordData?.data?.response?.resultUrl); + add(response.resultUrls); + add(response.resultUrl); // Check direct data fields - add(recordData?.data?.resultImageUrls); - add(recordData?.data?.resultImageUrl); - add(recordData?.data?.url); + add(data.resultImageUrls); + add(data.resultImageUrl); + add(data.url); return Array.from(urls); } @@ -352,31 +367,44 @@ async function handleKieImageGeneration({ const token = credentials?.apiKey || credentials?.accessToken; const timeoutMs = normalizePositiveNumber(body.timeout_ms, 300000); const pollIntervalMs = normalizePositiveNumber(body.poll_interval_ms, 2500); + const prompt = typeof body.prompt === "string" ? body.prompt : String(body.prompt ?? ""); + const size = typeof body.size === "string" ? body.size : undefined; + + if (!token) { + return saveImageErrorResult({ + provider, + model, + status: 401, + startTime, + error: "KIE API key is required", + }); + } // Check if model is a Market model (unified API) const fullRegistry = getImageProvider(provider); - const modelEntry = fullRegistry?.models?.find((m: any) => m.id === model); + const modelEntry = fullRegistry?.models?.find((m) => m.id === model); const isMarket = modelEntry?.isMarket || model.includes("/"); const { imageUrl } = extractImageInputs(body); let baseUrl = ""; - let payload: any = {}; + let payload: Record = {}; if (isMarket) { // Unified Market API endpoint baseUrl = `${providerConfig.baseUrl.replace(/\/$/, "")}/api/v1/jobs/createTask`; // Strip category prefix (e.g., "gpt/gpt-image-2" -> "gpt-image-2") const marketModelId = model.includes("/") ? model.split("/").pop() : model; - payload = { - model: marketModelId, - input: { - prompt: body.prompt, - aspect_ratio: mapImageSize(body.size, "1:1"), - }, + const input: Record = { + prompt, + aspect_ratio: mapImageSize(size, "1:1"), }; if (imageUrl) { - payload.input.image_url = imageUrl; + input.image_url = imageUrl; } + payload = { + model: marketModelId, + input, + }; } else { // Legacy/Direct endpoint const modelPath = model.replace("-t2i", "").replace("-i2i", ""); @@ -385,8 +413,8 @@ async function handleKieImageGeneration({ : `https://api.kie.ai/api/v1/${modelPath}/generate`; payload = { - prompt: body.prompt, - image_size: mapImageSize(body.size, "1:1"), + prompt, + image_size: mapImageSize(size, "1:1"), num_images: body.n || 1, }; } @@ -436,106 +464,58 @@ async function handleKieImageGeneration({ ? providerConfig.statusUrl : baseUrl.replace(/\/generate$/, "/record-info"); - const deadline = Date.now() + timeoutMs; - while (Date.now() < deadline) { - const pollUrl = new URL(statusUrl); - pollUrl.searchParams.set("taskId", String(taskId)); + const { data: recordData, state } = await kieExecutor.pollTask({ + statusUrl, + taskId: String(taskId), + token, + timeoutMs, + pollIntervalMs, + }); - const recordRes = await fetch(pollUrl.toString(), { - method: "GET", - headers: { - Authorization: `Bearer ${token}`, - }, + if (state === "success") { + if (log) { + log.info("IMAGE", `KIE poll success for task ${taskId}`); + } + const urls = normalizeKieImageResult(recordData); + const images = urls.map((url: string) => ({ url, revised_prompt: prompt })); + + return saveImageSuccessResult({ + provider, + model, + startTime, + requestBody: payload, + responseBody: { images_count: images.length }, + images, }); + } - if (!recordRes.ok) { - const errorText = await recordRes.text(); - return saveImageErrorResult({ - provider, - model, - status: recordRes.status, - startTime, - error: errorText, - requestBody: payload, - }); - } + const record = isJsonObject(recordData) ? recordData : {}; + const recordDataBody = isJsonObject(record.data) ? record.data : {}; + const errorMessage = + recordDataBody.errorMessage || + recordDataBody.failMsg || + record.msg || + "KIE image task failed"; - const recordData = await recordRes.json(); - const state = String( - recordData?.data?.status ?? - recordData?.data?.state ?? - recordData?.data?.successFlag ?? - recordData?.msg ?? - "PENDING" - ).toUpperCase(); - - if (state === "SUCCESS" || state === "1" || state === "FINISHED") { - if (log) { - log.info("IMAGE", `KIE poll success for task ${taskId}`); - } - const urls = normalizeKieImageResult(recordData); - const images = urls.map((url: string) => ({ url, revised_prompt: body.prompt })); - - return saveImageSuccessResult({ - provider, - model, - startTime, - requestBody: payload, - responseBody: { images_count: images.length }, - images, - }); - } - - // Expanded failure state detection - if ( - state === "FAIL" || - state === "FAILED" || - state === "ERROR" || - state === "2" || - state === "3" || - state.includes("FAIL") || - state.includes("ERROR") || - state === "CREATE_TASK_FAILED" || - state === "GENERATE_FAILED" - ) { - const errorMessage = - recordData?.data?.errorMessage || - recordData?.data?.failMsg || - recordData?.msg || - `KIE image task failed with status: ${state}`; - - if (log) { - log.error("IMAGE", `KIE poll failed for task ${taskId}: ${JSON.stringify(recordData)}`); - } - - return saveImageErrorResult({ - provider, - model, - status: 502, - startTime, - error: errorMessage, - requestBody: payload, - }); - } - - await sleep(pollIntervalMs); + if (log) { + log.error("IMAGE", `KIE poll failed for task ${taskId}: ${JSON.stringify(recordData)}`); } return saveImageErrorResult({ provider, model, - status: 504, + status: 502, startTime, - error: `KIE image polling timed out after ${timeoutMs}ms`, + error: String(errorMessage), requestBody: payload, }); - } catch (err) { + } catch (err: unknown) { return saveImageErrorResult({ provider, model, - status: 502, + status: getKieErrorStatus(err, 502), startTime, - error: `Image provider error: ${err instanceof Error ? err.message : String(err)}`, + error: `Image provider error: ${getKieErrorMessage(err, "KIE image generation failed")}`, }); } } diff --git a/open-sse/handlers/musicGeneration.ts b/open-sse/handlers/musicGeneration.ts index 0dc197231d..3a1202eeba 100644 --- a/open-sse/handlers/musicGeneration.ts +++ b/open-sse/handlers/musicGeneration.ts @@ -23,16 +23,7 @@ import { extractComfyOutputFiles, } from "../utils/comfyuiClient.ts"; import { saveCallLog } from "@/lib/usageDb"; -import { sleep } from "../utils/sleep.ts"; - -function getKieCallbackUrl(body: any): string { - return ( - body.callBackUrl || - body.callback_url || - body.callbackUrl || - "https://omniroute.local/api/kie/callback" - ); -} +import { getKieCallbackUrl, isJsonObject, parseKieResultJson } from "../utils/kieTask.ts"; function normalizeKieSunoModel(model: string): string { const map: Record = { @@ -42,42 +33,39 @@ function normalizeKieSunoModel(model: string): string { return map[model] || model; } -function parseKieResultJson(recordData: any): any { - try { - return typeof recordData?.data?.resultJson === "string" - ? JSON.parse(recordData.data.resultJson) - : recordData?.data?.resultJson || {}; - } catch { - return {}; - } -} - -function normalizeKieMusicTracks(recordData: any): any[] { +function normalizeKieMusicTracks(recordData: unknown): Array> { + const record = isJsonObject(recordData) ? recordData : {}; + const data = isJsonObject(record.data) ? record.data : {}; + const response = isJsonObject(data.response) ? data.response : {}; const resultJson = parseKieResultJson(recordData); const candidates = [ - recordData?.data?.response?.sunoData, - recordData?.data?.response?.data, - recordData?.data?.data, - recordData?.data?.sunoData, - resultJson?.sunoData, - resultJson?.data, - resultJson?.result, + response.sunoData, + response.data, + data.data, + data.sunoData, + resultJson.sunoData, + resultJson.data, + resultJson.result, ]; for (const candidate of candidates) { if (Array.isArray(candidate) && candidate.length > 0) { - return candidate; + return candidate + .map((track) => + isJsonObject(track) ? track : typeof track === "string" ? { audioUrl: track } : null + ) + .filter((track): track is Record => track !== null); } } const singleUrl = - recordData?.data?.response?.audioUrl || - recordData?.data?.response?.audio_url || - recordData?.data?.resultUrl || - recordData?.data?.audio_url || - resultJson?.audioUrl || - resultJson?.audio_url || - resultJson?.url; + response.audioUrl || + response.audio_url || + data.resultUrl || + data.audio_url || + resultJson.audioUrl || + resultJson.audio_url || + resultJson.url; return typeof singleUrl === "string" && singleUrl.length > 0 ? [{ audioUrl: singleUrl }] : []; } @@ -238,24 +226,42 @@ async function handleKieMusicGeneration({ }: { model: string; provider: string; - providerConfig: any; - body: any; - credentials: any; - log: any; + providerConfig: { + baseUrl: string; + statusUrl?: string; + }; + body: Record & { + prompt?: unknown; + timeout_ms?: unknown; + poll_interval_ms?: unknown; + }; + credentials?: { + apiKey?: string; + accessToken?: string; + } | null; + log?: { + info: (scope: string, message: string) => void; + error: (scope: string, message: string) => void; + } | null; }) { const startTime = Date.now(); const timeoutMs = Number(body.timeout_ms) > 0 ? Number(body.timeout_ms) : 300000; const pollIntervalMs = Number(body.poll_interval_ms) > 0 ? Number(body.poll_interval_ms) : 2500; const token = credentials?.apiKey || credentials?.accessToken; const baseUrl = providerConfig.baseUrl.replace(/\/$/, ""); + const prompt = typeof body.prompt === "string" ? body.prompt : String(body.prompt ?? ""); + + if (!token) { + return { success: false, status: 401, error: "KIE API key is required" }; + } // Check if model is a Market model const fullRegistry = getMusicProvider(provider); - const modelEntry = fullRegistry?.models?.find((m: any) => m.id === model); + const modelEntry = fullRegistry?.models?.find((m) => m.id === model); const isMarket = modelEntry?.isMarket || model.includes("/"); let url = ""; - let payload: any = {}; + let payload: Record = {}; if (isMarket) { url = `${baseUrl}/api/v1/jobs/createTask`; @@ -263,14 +269,14 @@ async function handleKieMusicGeneration({ model: model.includes("/") ? model.split("/").pop() : model, callBackUrl: getKieCallbackUrl(body), input: { - prompt: body.prompt, + prompt, instrumental: true, }, }; } else { url = `${baseUrl}/api/v1/generate`; payload = { - prompt: body.prompt, + prompt, customMode: false, instrumental: true, model: normalizeKieSunoModel(model), @@ -302,92 +308,58 @@ async function handleKieMusicGeneration({ return { success: false, status: 502, error: errorMessage }; } - const deadline = Date.now() + timeoutMs; const statusUrl = isMarket ? `${baseUrl}/api/v1/jobs/recordInfo` : providerConfig.statusUrl || `${baseUrl}/api/v1/generate/record-info`; - while (Date.now() < deadline) { - const pollUrl = new URL(statusUrl); - pollUrl.searchParams.set("taskId", String(taskId)); + const { data: recordData, state } = await kieExecutor.pollTask({ + statusUrl, + taskId: String(taskId), + token, + timeoutMs, + pollIntervalMs, + }); - const recordRes = await fetch(pollUrl.toString(), { - method: "GET", - headers: { - Authorization: `Bearer ${token}`, - }, - }); + if (state === "success") { + const tracks = normalizeKieMusicTracks(recordData); - if (!recordRes.ok) { - const errorText = await recordRes.text(); - return { success: false, status: recordRes.status, error: errorText }; - } + const audioFiles = tracks + .map((track) => + typeof track.audioUrl === "string" + ? track.audioUrl + : typeof track.audio_url === "string" + ? track.audio_url + : typeof track.url === "string" + ? track.url + : null + ) + .filter((url): url is string => typeof url === "string" && url.length > 0) + .map((url: string) => ({ url, format: "mp3" })); - const recordData = await recordRes.json(); - const state = String( - recordData?.data?.status ?? - recordData?.data?.state ?? - recordData?.data?.successFlag ?? - recordData?.msg ?? - "PENDING" - ).toUpperCase(); + saveCallLog({ + method: "POST", + path: "/v1/music/generations", + status: 200, + model: `${provider}/${model}`, + provider, + duration: Date.now() - startTime, + responseBody: { audio_count: audioFiles.length }, + }).catch(() => {}); - if (state === "SUCCESS" || state === "1" || state === "FINISHED") { - const tracks = normalizeKieMusicTracks(recordData); - - const audioFiles = tracks - .map((track: any) => { - return ( - typeof track?.audioUrl === "string" ? track.audioUrl : track?.audio_url || track?.url - ) as string; - }) - .filter((url: string) => typeof url === "string" && url.length > 0) - .map((url: string) => ({ url, format: "mp3" })); - - saveCallLog({ - method: "POST", - path: "/v1/music/generations", - status: 200, - model: `${provider}/${model}`, - provider, - duration: Date.now() - startTime, - responseBody: { audio_count: audioFiles.length }, - }).catch(() => {}); - - return { - success: true, - data: { created: Math.floor(Date.now() / 1000), data: audioFiles }, - }; - } - - if ( - state.includes("FAIL") || - state.includes("ERROR") || - state === "2" || - state === "3" || - state === "CREATE_TASK_FAILED" || - state === "GENERATE_AUDIO_FAILED" - ) { - const errorMessage = - recordData?.data?.errorMessage || - recordData?.data?.failMsg || - recordData?.msg || - `KIE music task failed with status: ${state}`; - return { success: false, status: 502, error: errorMessage }; - } - - await sleep(pollIntervalMs); + return { + success: true, + data: { created: Math.floor(Date.now() / 1000), data: audioFiles }, + }; } + const record = isJsonObject(recordData) ? recordData : {}; + const data = isJsonObject(record.data) ? record.data : {}; + const errorMessage = data.errorMessage || data.failMsg || record.msg || "KIE music task failed"; + return { success: false, status: 502, error: String(errorMessage) }; + } catch (err: unknown) { return { success: false, - status: 504, - error: `KIE music polling timed out after ${timeoutMs}ms`, - }; - } catch (err) { - return { - success: false, - status: 502, + status: isJsonObject(err) && Number.isFinite(Number(err.status)) ? Number(err.status) : 502, error: `Music provider error: ${err instanceof Error ? err.message : String(err)}`, }; } diff --git a/open-sse/handlers/videoGeneration.ts b/open-sse/handlers/videoGeneration.ts index 6ee997f275..25065c60f8 100644 --- a/open-sse/handlers/videoGeneration.ts +++ b/open-sse/handlers/videoGeneration.ts @@ -17,6 +17,7 @@ import { getVideoProvider, parseVideoModel } from "../config/videoRegistry.ts"; import { kieExecutor } from "../executors/kie.ts"; +import { isJsonObject, parseKieResultJson } from "../utils/kieTask.ts"; import { submitComfyWorkflow, pollComfyResult, @@ -24,7 +25,6 @@ import { extractComfyOutputFiles, } from "../utils/comfyuiClient.ts"; import { saveCallLog } from "@/lib/usageDb"; -import { sleep } from "../utils/sleep.ts"; /** * Handle video generation request @@ -269,23 +269,18 @@ async function handleSDWebUIVideoGeneration({ model, provider, providerConfig, b } } -function normalizeKieVideoResult(recordData: any): string[] { - let resultJson: Record = {}; - try { - resultJson = - typeof recordData?.data?.resultJson === "string" - ? JSON.parse(recordData.data.resultJson) - : recordData?.data?.resultJson || {}; - } catch { - resultJson = {}; - } +function normalizeKieVideoResult(recordData: unknown): string[] { + const record = isJsonObject(recordData) ? recordData : {}; + const data = isJsonObject(record.data) ? record.data : {}; + const response = isJsonObject(data.response) ? data.response : {}; + const resultJson = parseKieResultJson(recordData); const urls = Array.isArray(resultJson?.resultUrls) ? (resultJson.resultUrls as string[]) : Array.isArray(resultJson?.videoUrls) ? (resultJson.videoUrls as string[]) - : Array.isArray(recordData?.data?.response?.resultUrls) - ? (recordData.data.response.resultUrls as string[]) + : Array.isArray(response.resultUrls) + ? (response.resultUrls as string[]) : []; return urls.filter((url: unknown) => typeof url === "string" && url.length > 0); @@ -301,16 +296,37 @@ async function handleKieVideoGeneration({ }: { model: string; provider: string; - providerConfig: any; - body: any; - credentials: any; - log: any; + providerConfig: { + baseUrl: string; + statusUrl?: string; + }; + body: Record & { + prompt?: unknown; + duration?: unknown; + aspect_ratio?: unknown; + sound?: unknown; + timeout_ms?: unknown; + poll_interval_ms?: unknown; + }; + credentials?: { + apiKey?: string; + accessToken?: string; + } | null; + log?: { + info: (scope: string, message: string) => void; + error: (scope: string, message: string) => void; + } | null; }) { const startTime = Date.now(); const timeoutMs = Number(body.timeout_ms) > 0 ? Number(body.timeout_ms) : 300000; const pollIntervalMs = Number(body.poll_interval_ms) > 0 ? Number(body.poll_interval_ms) : 2500; const token = credentials?.apiKey || credentials?.accessToken; const baseUrl = providerConfig.baseUrl.replace(/\/$/, ""); + const prompt = typeof body.prompt === "string" ? body.prompt : String(body.prompt ?? ""); + + if (!token) { + return { success: false, status: 401, error: "KIE API key is required" }; + } // Strip category prefix (e.g., "veo/veo-3-1" -> "veo-3-1") const marketModelId = model.includes("/") ? model.split("/").pop() : model; @@ -318,9 +334,9 @@ async function handleKieVideoGeneration({ const payload = { model: marketModelId, input: { - prompt: body.prompt, + prompt, duration: body.duration ? String(body.duration) : "5", - aspect_ratio: body.aspect_ratio || "16:9", + aspect_ratio: typeof body.aspect_ratio === "string" ? body.aspect_ratio : "16:9", sound: body.sound === true, }, }; @@ -345,80 +361,44 @@ async function handleKieVideoGeneration({ return { success: false, status: 502, error: errorMessage }; } - const deadline = Date.now() + timeoutMs; const statusUrl = providerConfig.statusUrl || `${baseUrl}/api/v1/jobs/recordInfo`; - while (Date.now() < deadline) { - const pollUrl = new URL(statusUrl); - pollUrl.searchParams.set("taskId", String(taskId)); + const { data: recordData, state } = await kieExecutor.pollTask({ + statusUrl, + taskId: String(taskId), + token, + timeoutMs, + pollIntervalMs, + }); - const recordRes = await fetch(pollUrl.toString(), { - method: "GET", - headers: { - Authorization: `Bearer ${token}`, - }, - }); + if (state === "success") { + const videoUrls = normalizeKieVideoResult(recordData); + const videos = videoUrls.map((url) => ({ url, format: "mp4" })); - if (!recordRes.ok) { - const errorText = await recordRes.text(); - return { success: false, status: recordRes.status, error: errorText }; - } + saveCallLog({ + method: "POST", + path: "/v1/videos/generations", + status: 200, + model: `${provider}/${model}`, + provider, + duration: Date.now() - startTime, + responseBody: { videos_count: videos.length }, + }).catch(() => {}); - const recordData = await recordRes.json(); - const state = String( - recordData?.data?.state || recordData?.data?.status || "generating" - ).toLowerCase(); - - if (state === "success" || state === "1" || state === "finished") { - const videoUrls = normalizeKieVideoResult(recordData); - const videos = videoUrls.map((url) => ({ url, format: "mp4" })); - - saveCallLog({ - method: "POST", - path: "/v1/videos/generations", - status: 200, - model: `${provider}/${model}`, - provider, - duration: Date.now() - startTime, - responseBody: { videos_count: videos.length }, - }).catch(() => {}); - - return { - success: true, - data: { created: Math.floor(Date.now() / 1000), data: videos }, - }; - } - - if ( - state === "fail" || - state === "failed" || - state === "error" || - state === "2" || - state === "3" || - state.includes("fail") || - state.includes("error") || - state.includes("failed") - ) { - const errorMessage = - recordData?.data?.failMsg || - recordData?.data?.errorMessage || - recordData?.msg || - `KIE video task failed with state: ${state}`; - return { success: false, status: 502, error: errorMessage }; - } - - await sleep(pollIntervalMs); + return { + success: true, + data: { created: Math.floor(Date.now() / 1000), data: videos }, + }; } + const record = isJsonObject(recordData) ? recordData : {}; + const data = isJsonObject(record.data) ? record.data : {}; + const errorMessage = data.failMsg || data.errorMessage || record.msg || "KIE video task failed"; + return { success: false, status: 502, error: String(errorMessage) }; + } catch (err: unknown) { return { success: false, - status: 504, - error: `KIE video polling timed out after ${timeoutMs}ms`, - }; - } catch (err) { - return { - success: false, - status: 502, + status: isJsonObject(err) && Number.isFinite(Number(err.status)) ? Number(err.status) : 502, error: `Video provider error: ${err instanceof Error ? err.message : String(err)}`, }; } diff --git a/open-sse/utils/kieTask.ts b/open-sse/utils/kieTask.ts new file mode 100644 index 0000000000..78f52fd59c --- /dev/null +++ b/open-sse/utils/kieTask.ts @@ -0,0 +1,97 @@ +export type JsonObject = Record; + +export type KieTaskState = "success" | "failed" | "pending"; + +export type KieCallbackBody = { + callBackUrl?: unknown; + callback_url?: unknown; + callbackUrl?: unknown; +}; + +export function isJsonObject(value: unknown): value is JsonObject { + return typeof value === "object" && value !== null && !Array.isArray(value); +} + +export function getKieCallbackUrl(body: KieCallbackBody = {}): string { + const callbackUrl = body.callBackUrl ?? body.callback_url ?? body.callbackUrl; + return typeof callbackUrl === "string" && callbackUrl.trim().length > 0 + ? callbackUrl + : "https://omniroute.local/api/kie/callback"; +} + +export function parseKieResultJson(recordData: unknown): JsonObject { + const data = isJsonObject(recordData) && isJsonObject(recordData.data) ? recordData.data : {}; + const resultJson = data.resultJson; + + if (typeof resultJson === "string") { + try { + const parsed = JSON.parse(resultJson) as unknown; + return isJsonObject(parsed) ? parsed : {}; + } catch { + return {}; + } + } + + return isJsonObject(resultJson) ? resultJson : {}; +} + +export function normalizeKieTaskState(recordData: unknown): KieTaskState { + const record = isJsonObject(recordData) ? recordData : {}; + const data = isJsonObject(record.data) ? record.data : {}; + const state = String( + data.status ?? data.state ?? data.successFlag ?? record.msg ?? "PENDING" + ).toUpperCase(); + + if ( + state === "SUCCESS" || + state === "1" || + state === "FINISHED" || + state === "COMPLETE" || + state === "COMPLETED" || + state === "FIRST_SUCCESS" || + state === "ALL_SUCCESS" || + state.includes("SUCCESS") + ) { + return "success"; + } + + if ( + state === "FAIL" || + state === "FAILED" || + state === "ERROR" || + state === "2" || + state === "3" || + state.includes("FAIL") || + state.includes("ERROR") || + state === "CREATE_TASK_FAILED" || + state === "GENERATE_FAILED" || + state === "GENERATE_AUDIO_FAILED" + ) { + return "failed"; + } + + return "pending"; +} + +export function getKieErrorStatus(error: unknown, fallback = 502): number { + if (isJsonObject(error)) { + const status = Number(error.status); + if (Number.isFinite(status) && status > 0) { + return status; + } + } + + return fallback; +} + +export function getKieErrorMessage(error: unknown, fallback: string): string { + if (error instanceof Error && error.message) { + return error.message; + } + + if (isJsonObject(error) && typeof error.message === "string" && error.message.length > 0) { + return error.message; + } + + return typeof error === "string" && error.length > 0 ? error : fallback; +} diff --git a/src/lib/providers/validation.ts b/src/lib/providers/validation.ts index 67b88470f6..3cf56f0a6c 100644 --- a/src/lib/providers/validation.ts +++ b/src/lib/providers/validation.ts @@ -279,7 +279,7 @@ async function validateDirectChatProvider({ url, headers, body, providerSpecific } return { valid: false, error: `Validation failed: ${response.status}` }; - } catch (error: any) { + } catch (error: unknown) { return toValidationErrorResult(error); } } @@ -577,7 +577,7 @@ async function validateKieProvider({ apiKey, providerSpecificData = {} }: any) { } return { valid: false, error: `Validation failed: ${chatRes.status}` }; - } catch (error: any) { + } catch (error: unknown) { return toValidationErrorResult(error); } }