UNPKG

@huggingface/inference

Version:

Typescript client for the Hugging Face Inference Providers and Inference Endpoints

153 lines (152 loc) 7.09 kB
"use strict"; Object.defineProperty(exports, "__esModule", { value: true }); exports.ZaiImageToTextTask = exports.ZaiTextToImageTask = exports.ZaiConversationalTask = void 0; const errors_js_1 = require("../errors.js"); const isUrl_js_1 = require("../lib/isUrl.js"); const base64FromBytes_js_1 = require("../utils/base64FromBytes.js"); const dataUrlFromBlob_js_1 = require("../utils/dataUrlFromBlob.js"); const delay_js_1 = require("../utils/delay.js"); const omit_js_1 = require("../utils/omit.js"); const providerHelper_js_1 = require("./providerHelper.js"); const ZAI_API_BASE_URL = "https://api.z.ai"; class ZaiTask extends providerHelper_js_1.TaskProviderHelper { constructor() { super("zai-org", ZAI_API_BASE_URL); } prepareHeaders(params, binary) { const headers = super.prepareHeaders(params, binary); headers["x-source-channel"] = "hugging_face"; headers["accept-language"] = "en-US,en"; return headers; } } class ZaiConversationalTask extends providerHelper_js_1.BaseConversationalTask { constructor() { super("zai-org", ZAI_API_BASE_URL); } prepareHeaders(params, binary) { const headers = super.prepareHeaders(params, binary); headers["x-source-channel"] = "hugging_face"; headers["accept-language"] = "en-US,en"; return headers; } makeRoute() { return "/api/paas/v4/chat/completions"; } } exports.ZaiConversationalTask = ZaiConversationalTask; const MAX_POLL_ATTEMPTS = 60; const POLL_INTERVAL_MS = 5000; class ZaiTextToImageTask extends ZaiTask { makeRoute() { return "/api/paas/v4/async/images/generations"; } preparePayload(params) { return { ...(0, omit_js_1.omit)(params.args, ["inputs", "parameters"]), ...params.args.parameters, model: params.model, prompt: params.args.inputs, }; } async getResponse(response, url, headers, outputType, signal) { if (!url || !headers) { throw new errors_js_1.InferenceClientInputError(`URL and headers are required for 'text-to-image' task`); } if (typeof response !== "object" || !response || !("task_status" in response) || !("id" in response) || typeof response.id !== "string") { throw new errors_js_1.InferenceClientProviderOutputError(`Received malformed response from ZAI text-to-image API: expected { id: string, task_status: string }, got: ${JSON.stringify(response)}`); } if (response.task_status === "FAIL") { throw new errors_js_1.InferenceClientProviderOutputError("ZAI API returned task status: FAIL"); } const taskId = response.id; const parsedUrl = new URL(url); const baseUrl = `${parsedUrl.protocol}//${parsedUrl.host}${parsedUrl.host === "router.huggingface.co" ? "/zai-org" : ""}`; const pollUrl = `${baseUrl}/api/paas/v4/async-result/${taskId}`; const pollHeaders = { ...headers, "x-source-channel": "hugging_face", "accept-language": "en-US,en", }; for (let attempt = 0; attempt < MAX_POLL_ATTEMPTS; attempt++) { await (0, delay_js_1.delay)(POLL_INTERVAL_MS, signal); const resp = await fetch(pollUrl, { method: "GET", headers: pollHeaders, signal, }); if (!resp.ok) { throw new errors_js_1.InferenceClientProviderApiError(`Failed to fetch result from ZAI text-to-image API: ${resp.status}`, { url: pollUrl, method: "GET" }, { requestId: resp.headers.get("x-request-id") ?? "", status: resp.status, body: await resp.text() }); } const result = await resp.json(); if (result.task_status === "FAIL") { throw new errors_js_1.InferenceClientProviderOutputError("ZAI text-to-image API task failed"); } if (result.task_status === "SUCCESS") { if (!result.image_result || !Array.isArray(result.image_result) || result.image_result.length === 0 || typeof result.image_result[0]?.url !== "string" || !(0, isUrl_js_1.isUrl)(result.image_result[0].url)) { throw new errors_js_1.InferenceClientProviderOutputError(`Received malformed response from ZAI text-to-image API: expected { image_result: Array<{ url: string }> }, got: ${JSON.stringify(result)}`); } const imageUrl = result.image_result[0].url; if (outputType === "json") { return { ...result }; } if (outputType === "url") { return imageUrl; } const imageResponse = await fetch(imageUrl, { signal }); const blob = await imageResponse.blob(); return outputType === "dataUrl" ? (0, dataUrlFromBlob_js_1.dataUrlFromBlob)(blob) : blob; } } throw new errors_js_1.InferenceClientProviderOutputError(`Timed out while waiting for the result from ZAI API - aborting after ${MAX_POLL_ATTEMPTS} attempts`); } } exports.ZaiTextToImageTask = ZaiTextToImageTask; class ZaiImageToTextTask extends ZaiTask { makeRoute() { return "/api/paas/v4/layout_parsing"; } async preparePayloadAsync(args, signal) { const blob = "data" in args && args.data instanceof Blob ? args.data : "inputs" in args ? typeof args.inputs === "string" && (0, isUrl_js_1.isUrl)(args.inputs) ? await fetch(args.inputs, { signal }).then((r) => r.blob()) : args.inputs instanceof Blob ? args.inputs : undefined : undefined; if (!blob || !(blob instanceof Blob)) { throw new errors_js_1.InferenceClientInputError("ZAI image-to-text requires a URL string or Blob as inputs"); } const mimeType = blob.type || "image/png"; const b64 = (0, base64FromBytes_js_1.base64FromBytes)(new Uint8Array(await blob.arrayBuffer())); const file = `data:${mimeType};base64,${b64}`; return { ...("data" in args ? (0, omit_js_1.omit)(args, "data") : (0, omit_js_1.omit)(args, "inputs")), inputs: file, }; } preparePayload(params) { return { model: params.model, file: params.args.inputs, }; } async getResponse(response) { const mdResults = response?.md_results; if (typeof mdResults !== "string") { throw new errors_js_1.InferenceClientProviderOutputError(`Received malformed response from ZAI layout_parsing API: expected { md_results: string }, got: ${JSON.stringify(response)}`); } return { generated_text: mdResults, generatedText: mdResults }; } } exports.ZaiImageToTextTask = ZaiImageToTextTask;