cactus-ai-sdk-provider
Version:
Cactus provider for Vercel AI SDK - enables on-device LLM inference in React Native
392 lines (386 loc) • 11.6 kB
JavaScript
"use strict";
var __defProp = Object.defineProperty;
var __getOwnPropDesc = Object.getOwnPropertyDescriptor;
var __getOwnPropNames = Object.getOwnPropertyNames;
var __hasOwnProp = Object.prototype.hasOwnProperty;
var __export = (target, all) => {
for (var name in all)
__defProp(target, name, { get: all[name], enumerable: true });
};
var __copyProps = (to, from, except, desc) => {
if (from && typeof from === "object" || typeof from === "function") {
for (let key of __getOwnPropNames(from))
if (!__hasOwnProp.call(to, key) && key !== except)
__defProp(to, key, { get: () => from[key], enumerable: !(desc = __getOwnPropDesc(from, key)) || desc.enumerable });
}
return to;
};
var __toCommonJS = (mod) => __copyProps(__defProp({}, "__esModule", { value: true }), mod);
// index.ts
var index_exports = {};
__export(index_exports, {
CactusChatLanguageModel: () => CactusChatLanguageModel,
ModelManager: () => ModelManager,
ModelStatus: () => ModelStatus,
cactus: () => cactus,
createCactus: () => createCactus
});
module.exports = __toCommonJS(index_exports);
// provider.ts
var import_provider_utils = require("@ai-sdk/provider-utils");
// cactus-chat-language-model.ts
var import_cactus_react_native = require("cactus-react-native");
// model-manager.ts
var ModelManager = class {
rnfs = null;
cache = /* @__PURE__ */ new Map();
constructor() {
try {
this.rnfs = require("react-native-fs");
} catch (e) {
}
}
/**
* Check if RNFS is available for model management features
*/
get isRNFSAvailable() {
return this.rnfs !== null;
}
/**
* Get the local file path for a model URL
* @throws Error if RNFS is not available
*/
getModelPath(urlOrPath) {
if (!urlOrPath.startsWith("http://") && !urlOrPath.startsWith("https://")) {
return urlOrPath;
}
if (!this.rnfs) {
throw new Error(
"react-native-fs is required for automatic model paths. Please install it or provide a direct file path instead of a URL."
);
}
const filename = urlOrPath.split("/").pop()?.replace(/[^a-zA-Z0-9.-]/g, "_") || "model.gguf";
return `${this.rnfs.DocumentDirectoryPath}/cactus-models/${filename}`;
}
/**
* Download a model from a URL
* @throws Error if RNFS is not available
*/
async downloadModel(url, options) {
if (!this.rnfs) {
throw new Error(
"react-native-fs is required for model downloads. Please install it with: npm install react-native-fs"
);
}
const localPath = this.getModelPath(url);
const modelDir = `${this.rnfs.DocumentDirectoryPath}/cactus-models`;
if (!await this.rnfs.exists(modelDir)) {
await this.rnfs.mkdir(modelDir);
}
if (await this.rnfs.exists(localPath)) {
this.cache.set(url, { url, localPath, status: "ready" });
return localPath;
}
this.cache.set(url, { url, localPath, status: "pending" });
try {
await this.rnfs.downloadFile({
fromUrl: url,
toFile: localPath,
progress: (res) => {
if (options?.onProgress) {
options.onProgress(res.bytesWritten / res.contentLength);
}
}
}).promise;
this.cache.set(url, { url, localPath, status: "ready" });
return localPath;
} catch (error) {
this.cache.set(url, { url, localPath, status: "error" });
try {
if (await this.rnfs.exists(localPath)) {
await this.rnfs.unlink(localPath);
}
} catch {
}
throw error;
}
}
/**
* Check if a model is downloaded
*/
async isDownloaded(urlOrPath) {
if (!urlOrPath.startsWith("http://") && !urlOrPath.startsWith("https://")) {
return true;
}
const cached = this.cache.get(urlOrPath);
if (cached?.status === "ready") {
return true;
}
if (!this.rnfs) {
return false;
}
try {
const localPath = this.getModelPath(urlOrPath);
const exists = await this.rnfs.exists(localPath);
if (exists) {
this.cache.set(urlOrPath, { url: urlOrPath, localPath, status: "ready" });
}
return exists;
} catch {
return false;
}
}
/**
* Delete a downloaded model
*/
async deleteModel(urlOrPath) {
if (!this.rnfs) {
throw new Error("react-native-fs is required to delete models");
}
const localPath = this.getModelPath(urlOrPath);
if (await this.rnfs.exists(localPath)) {
await this.rnfs.unlink(localPath);
}
this.cache.delete(urlOrPath);
}
/**
* List all downloaded models
*/
async listModels() {
if (!this.rnfs) {
return Array.from(this.cache.values());
}
const modelDir = `${this.rnfs.DocumentDirectoryPath}/cactus-models`;
try {
if (!await this.rnfs.exists(modelDir)) {
return [];
}
const files = await this.rnfs.readDir(modelDir);
const models = [];
for (const file of files) {
if (file.isFile() && file.name.endsWith(".gguf")) {
let found = false;
for (const [url, entry] of this.cache.entries()) {
if (entry.localPath === file.path) {
models.push(entry);
found = true;
break;
}
}
if (!found) {
models.push({
url: file.path,
localPath: file.path,
status: "ready"
});
}
}
}
return models;
} catch {
return [];
}
}
/**
* Get storage info for all models
*/
async getStorageInfo() {
if (!this.rnfs) {
return { totalSize: 0, modelCount: 0 };
}
const models = await this.listModels();
let totalSize = 0;
for (const model of models) {
try {
const stat = await this.rnfs.stat(model.localPath);
totalSize += stat.size;
} catch {
}
}
return {
totalSize,
modelCount: models.length
};
}
};
// cactus-chat-language-model.ts
var ModelStatus = /* @__PURE__ */ ((ModelStatus2) => {
ModelStatus2[ModelStatus2["IDLE"] = 0] = "IDLE";
ModelStatus2[ModelStatus2["DOWNLOADING"] = 1] = "DOWNLOADING";
ModelStatus2[ModelStatus2["INITIALIZING"] = 2] = "INITIALIZING";
ModelStatus2[ModelStatus2["READY"] = 3] = "READY";
ModelStatus2[ModelStatus2["ERROR"] = 4] = "ERROR";
return ModelStatus2;
})(ModelStatus || {});
var CactusChatLanguageModel = class {
constructor(modelUrlOrPath, settings, config) {
this.settings = settings;
this.config = config;
this.modelUrl = modelUrlOrPath;
try {
this.modelId = this.modelManager.getModelPath(modelUrlOrPath);
} catch {
this.modelId = modelUrlOrPath;
}
this.provider = config.provider;
}
specificationVersion = "v2";
provider;
modelId;
modelUrl;
lm = null;
status = 0 /* IDLE */;
lastError = null;
modelManager = new ModelManager();
getStatus() {
return this.status;
}
getLastError() {
return this.lastError;
}
/**
* List all downloaded models
*/
async listDownloadedModels() {
return this.modelManager.listModels();
}
/**
* Get storage info for all downloaded models
*/
async getStorageInfo() {
return this.modelManager.getStorageInfo();
}
/**
* Download the model if it's a URL
*/
async downloadModel(options = {}) {
this.status = 1 /* DOWNLOADING */;
try {
const localPath = await this.modelManager.downloadModel(this.modelUrl, options);
this.status = 0 /* IDLE */;
return localPath;
} catch (e) {
this.lastError = e;
this.status = 4 /* ERROR */;
throw e;
}
}
/**
* Delete the downloaded model
*/
async deleteModel() {
await this.modelManager.deleteModel(this.modelUrl);
this.status = 0 /* IDLE */;
}
async initialize() {
if (this.status === 3 /* READY */) return;
const isDownloaded = await this.modelManager.isDownloaded(this.modelUrl);
if (!isDownloaded) {
throw new Error("Model not downloaded. Call downloadModel() first.");
}
this.status = 2 /* INITIALIZING */;
try {
const { lm, error } = await import_cactus_react_native.CactusLM.init({ model: this.modelId });
if (error) throw error;
this.lm = lm;
this.status = 3 /* READY */;
} catch (e) {
this.lastError = e;
this.status = 4 /* ERROR */;
throw e;
}
}
assertIsReady() {
if (this.status !== 3 /* READY */ || !this.lm) {
throw new Error(
`Model not ready. Status: ${ModelStatus[this.status]}. Error: ${this.lastError?.message}`
);
}
}
convertToCactusMessages(prompt) {
return prompt.map((message) => {
if (message.role === "system") {
return { role: "system", content: message.content };
}
const content = message.content.filter((part) => part.type === "text").map((part) => part.text).join("");
return { role: message.role, content };
});
}
getCactusParams(options) {
const params = {};
if (options.temperature != null) params.temperature = options.temperature;
if (options.maxOutputTokens != null) params.n_predict = options.maxOutputTokens;
if (options.stopSequences != null) params.stop = options.stopSequences;
return params;
}
get supportedUrls() {
return {};
}
// cactus is on-device only
async doGenerate(options) {
this.assertIsReady();
const messages = this.convertToCactusMessages(options.prompt);
const params = this.getCactusParams(options);
const result = await this.lm.completion(
messages,
params
);
return {
content: [{ type: "text", text: result.content }],
finishReason: "stop",
usage: {
inputTokens: result.tokens_evaluated,
outputTokens: result.tokens_predicted,
totalTokens: result.tokens_evaluated + result.tokens_predicted
},
warnings: []
};
}
async doStream(options) {
this.assertIsReady();
const messages = this.convertToCactusMessages(options.prompt);
const params = this.getCactusParams(options);
const stream = new ReadableStream({
start: async (controller) => {
controller.enqueue({ type: "stream-start", warnings: [] });
try {
const result = await this.lm.completion(messages, params, (data) => {
controller.enqueue({ type: "text-delta", id: "text", delta: data.token });
});
controller.enqueue({
type: "finish",
finishReason: "stop",
usage: {
inputTokens: result.tokens_evaluated,
outputTokens: result.tokens_predicted,
totalTokens: result.tokens_evaluated + result.tokens_predicted
}
});
controller.close();
} catch (error) {
controller.error(error);
}
}
});
return { stream, warnings: [] };
}
};
// provider.ts
function createCactus(options = {}) {
const createChatModel = (modelUrl, settings = {}) => new CactusChatLanguageModel(modelUrl, settings, {
provider: "cactus",
generateId: options.generateId ?? import_provider_utils.generateId
});
const provider = function(modelUrl, settings) {
if (new.target) {
throw new Error(
"The model factory function cannot be called with the new keyword."
);
}
return createChatModel(modelUrl, settings);
};
provider.languageModel = createChatModel;
return provider;
}
var cactus = createCactus();
//# sourceMappingURL=index.cjs.map