#![deny(clippy::all)]
use napi::bindgen_prelude::*;
use napi::Task;
use napi_derive::napi;
use ruvector_typesafe_core::engine::optimize::CampaignSpec;
use ruvector_typesafe_core::engine::{
Engine as CoreEngine, EngineOptions as CoreEngineOptions, LabeledExample,
};
use ruvector_typesafe_core::hash_embedder::HashEmbedder;
use ruvector_typesafe_core::{DecisionRequest, Embedder, TypesafeError};
use serde::Deserialize;
use std::sync::{Arc, RwLock};
type BoxedEngine = CoreEngine<Box<dyn Embedder>>;
type SharedEngine = Arc<RwLock<BoxedEngine>>;
#[napi]
pub fn version() -> String {
env!("CARGO_PKG_VERSION").to_string()
}
#[derive(Deserialize)]
struct EngineOptions {
embedder: EmbedderSpec,
#[serde(default = "default_dims")]
dims: usize,
#[serde(default)]
engine: Option<serde_json::Value>,
}
#[derive(Deserialize)]
#[serde(untagged)]
enum EmbedderSpec {
Named(String),
Kinded {
kind: String,
#[serde(rename = "modelDir", default)]
model_dir: Option<String>,
#[serde(default)]
manifest: Option<String>,
#[serde(default)]
model: Option<String>,
},
}
#[cfg(feature = "native-onnx")]
fn select_manifest(
json: &str,
model: Option<&str>,
dir: &str,
) -> std::result::Result<ruvector_embed_core::ModelManifest, String> {
use ruvector_embed_core::{ManifestFile, ModelManifest};
if let Ok(m) = ModelManifest::from_json(json) {
return Ok(m);
}
let file = ManifestFile::from_json(json).map_err(|e| format!("manifest: {e}"))?;
let base = std::path::Path::new(dir)
.file_name()
.and_then(|s| s.to_str())
.unwrap_or("");
let want = model.unwrap_or(base);
file.get(want)
.or_else(|| file.models.first())
.cloned()
.ok_or_else(|| format!("manifest: no entry named {want:?} and the collection is empty"))
}
fn default_dims() -> usize {
256
}
#[derive(Deserialize)]
struct TrainInput {
question: String,
#[serde(default)]
examples: Vec<ExampleInput>,
}
#[derive(Deserialize)]
struct ExampleInput {
text: String,
label: String,
}
fn make_engine(options_json: &str) -> std::result::Result<BoxedEngine, String> {
let opts: EngineOptions =
serde_json::from_str(options_json).map_err(|e| format!("invalid options JSON: {e}"))?;
let core_opts =
CoreEngineOptions::from_json_opt(opts.engine.as_ref()).map_err(|e| e.to_string())?;
Ok(CoreEngine::with_options(build_embedder(&opts)?, core_opts))
}
fn build_embedder(opts: &EngineOptions) -> std::result::Result<Box<dyn Embedder>, String> {
match &opts.embedder {
EmbedderSpec::Named(s) if s == "hash" => Ok(Box::new(HashEmbedder::new(opts.dims))),
EmbedderSpec::Named(s) => Err(format!("unknown embedder \"{s}\"")),
EmbedderSpec::Kinded {
kind,
model_dir,
manifest,
model,
} if kind == "onnx" => {
build_onnx(model_dir.as_deref(), manifest.as_deref(), model.as_deref())
}
EmbedderSpec::Kinded { kind, .. } => Err(format!("unknown embedder kind \"{kind}\"")),
}
}
#[cfg(feature = "native-onnx")]
fn build_onnx(
model_dir: Option<&str>,
manifest: Option<&str>,
model: Option<&str>,
) -> std::result::Result<Box<dyn Embedder>, String> {
use ruvector_embed_core::OrtEmbedder;
let dir = model_dir.ok_or("onnx embedder requires \"modelDir\"")?;
let manifest_src = manifest.ok_or("onnx embedder requires \"manifest\"")?;
let manifest_json =
std::fs::read_to_string(manifest_src).unwrap_or_else(|_| manifest_src.to_string());
let manifest = select_manifest(&manifest_json, model, dir)?;
let embedder =
OrtEmbedder::from_manifest(dir, &manifest).map_err(|e| format!("onnx load: {e}"))?;
Ok(Box::new(embedder))
}
#[cfg(not(feature = "native-onnx"))]
fn build_onnx(
_model_dir: Option<&str>,
_manifest: Option<&str>,
_model: Option<&str>,
) -> std::result::Result<Box<dyn Embedder>, String> {
Err("onnx embedder backend is not built into this binary \
(rebuild with --features native-onnx)"
.to_string())
}
#[napi]
pub struct Engine {
inner: SharedEngine,
}
#[napi]
impl Engine {
#[napi(constructor)]
pub fn new(options_json: String) -> Result<Self> {
let engine = make_engine(&options_json).map_err(Error::from_reason)?;
Ok(Self {
inner: Arc::new(RwLock::new(engine)),
})
}
#[napi(js_name = "decideJson")]
pub fn decide_json(&self, request_json: String) -> String {
let guard = self.inner.read().unwrap_or_else(|p| p.into_inner());
decide_to_json(&guard, &request_json)
}
#[napi(ts_return_type = "Promise<string>")]
pub fn decide(&self, request_json: String) -> AsyncTask<DecideTask> {
AsyncTask::new(DecideTask {
engine: self.inner.clone(),
request_json,
})
}
#[napi(js_name = "trainJson")]
pub fn train_json(&self, train_json: String) -> String {
let input: TrainInput = match serde_json::from_str(&train_json) {
Ok(v) => v,
Err(e) => return invalid_json(&format!("train JSON parse error: {e}")),
};
let examples: Vec<LabeledExample> = input
.examples
.into_iter()
.map(|e| LabeledExample {
text: e.text,
label: e.label,
})
.collect();
let mut guard = self.inner.write().unwrap_or_else(|p| p.into_inner());
match guard.train(&input.question, &examples) {
Ok(report) => {
serde_json::to_string(&report).unwrap_or_else(|e| embedder_json(&e.to_string()))
}
Err(e) => error_json(&e),
}
}
#[napi(js_name = "statsJson")]
pub fn stats_json(&self) -> String {
let guard = self.inner.read().unwrap_or_else(|p| p.into_inner());
let embedder = guard.embedder();
let summary = guard.bank_summary();
serde_json::json!({
"embedderId": embedder.id(),
"dims": embedder.dims(),
"questionsCompiled": 0,
"examples": summary.total,
"promotable": summary.promotable,
})
.to_string()
}
#[napi(js_name = "optimizeJson")]
pub fn optimize_json(&self, campaign_json: String) -> String {
let spec: CampaignSpec = match serde_json::from_str(&campaign_json) {
Ok(s) => s,
Err(e) => return invalid_json(&format!("campaign JSON parse error: {e}")),
};
let guard = self.inner.read().unwrap_or_else(|p| p.into_inner());
match guard.optimize(&spec) {
Ok(report) => {
serde_json::to_string(&report).unwrap_or_else(|e| embedder_json(&e.to_string()))
}
Err(e) => error_json(&e),
}
}
#[napi(js_name = "exportBankJson")]
pub fn export_bank_json(&self) -> String {
let guard = self.inner.read().unwrap_or_else(|p| p.into_inner());
match guard.export_bank() {
Ok(s) => s,
Err(e) => error_json(&e),
}
}
#[napi(js_name = "importBankJson")]
pub fn import_bank_json(&self, bank_json: String) -> String {
let mut guard = self.inner.write().unwrap_or_else(|p| p.into_inner());
match guard.import_bank(&bank_json) {
Ok(()) => "{\"ok\":true}".to_string(),
Err(e) => error_json(&e),
}
}
}
pub struct DecideTask {
engine: SharedEngine,
request_json: String,
}
impl Task for DecideTask {
type Output = String;
type JsValue = String;
fn compute(&mut self) -> Result<Self::Output> {
let guard = self.engine.read().unwrap_or_else(|p| p.into_inner());
Ok(decide_to_json(&guard, &self.request_json))
}
fn resolve(&mut self, _env: Env, output: Self::Output) -> Result<Self::JsValue> {
Ok(output)
}
}
fn decide_to_json(engine: &BoxedEngine, request_json: &str) -> String {
let req: DecisionRequest = match serde_json::from_str(request_json) {
Ok(r) => r,
Err(e) => return invalid_json(&format!("request JSON parse error: {e}")),
};
match engine.decide(&req) {
Ok(resp) => serde_json::to_string(&resp).unwrap_or_else(|e| embedder_json(&e.to_string())),
Err(e) => error_json(&e),
}
}
fn error_json(e: &TypesafeError) -> String {
let (kind, message) = match e {
TypesafeError::Limit(m) => ("limit", (*m).to_string()),
TypesafeError::Invalid(m) => ("invalid", m.clone()),
TypesafeError::Embedder(m) => ("embedder", m.clone()),
};
serde_json::json!({ "error": { "kind": kind, "message": message } }).to_string()
}
fn invalid_json(message: &str) -> String {
serde_json::json!({ "error": { "kind": "invalid", "message": message } }).to_string()
}
fn embedder_json(message: &str) -> String {
serde_json::json!({ "error": { "kind": "embedder", "message": message } }).to_string()
}