@huggingface/inference
Version:
Typescript client for the Hugging Face Inference Providers and Inference Endpoints
153 lines (152 loc) • 7.09 kB
JavaScript
;
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;