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 wasm_bindgen::prelude::*;
type BoxedEngine = CoreEngine<Box<dyn Embedder>>;
#[wasm_bindgen]
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 },
}
fn default_dims() -> usize {
256
}
#[derive(Deserialize)]
struct TrainInput {
question: String,
#[serde(default)]
examples: Vec<ExampleInput>,
}
#[derive(Deserialize)]
struct ExampleInput {
text: String,
label: String,
}
#[wasm_bindgen]
pub struct Engine {
inner: BoxedEngine,
}
#[wasm_bindgen]
impl Engine {
#[wasm_bindgen(constructor)]
pub fn new(options_json: &str) -> std::result::Result<Engine, JsValue> {
let opts: EngineOptions = serde_json::from_str(options_json)
.map_err(|e| JsValue::from_str(&format!("invalid options JSON: {e}")))?;
let core_opts = CoreEngineOptions::from_json_opt(opts.engine.as_ref())
.map_err(|e| JsValue::from_str(&e.to_string()))?;
match opts.embedder {
EmbedderSpec::Named(ref s) if s == "hash" => Ok(Engine {
inner: CoreEngine::with_options(Box::new(HashEmbedder::new(opts.dims)), core_opts),
}),
EmbedderSpec::Kinded { ref kind } if kind == "onnx" => Err(JsValue::from_str(
"onnx embedder requires Engine.fromBytes(optionsJson, modelBytes, tokenizerBytes)",
)),
EmbedderSpec::Named(s) => Err(JsValue::from_str(&format!("unknown embedder \"{s}\""))),
EmbedderSpec::Kinded { kind } => Err(JsValue::from_str(&format!(
"unknown embedder kind \"{kind}\""
))),
}
}
#[wasm_bindgen(js_name = fromBytes)]
pub fn from_bytes(
options_json: &str,
model_bytes: &[u8],
tokenizer_bytes: &[u8],
) -> std::result::Result<Engine, JsValue> {
build_onnx(options_json, model_bytes, tokenizer_bytes)
}
#[wasm_bindgen(js_name = decideJson)]
pub fn decide_json(&self, 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 self.inner.decide(&req) {
Ok(resp) => {
serde_json::to_string(&resp).unwrap_or_else(|e| embedder_json(&e.to_string()))
}
Err(e) => error_json(&e),
}
}
#[wasm_bindgen(js_name = trainJson)]
pub fn train_json(&mut self, train_json: &str) -> 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();
match self.inner.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),
}
}
#[wasm_bindgen(js_name = statsJson)]
pub fn stats_json(&self) -> String {
let embedder = self.inner.embedder();
let summary = self.inner.bank_summary();
serde_json::json!({
"embedderId": embedder.id(),
"dims": embedder.dims(),
"questionsCompiled": 0,
"examples": summary.total,
"promotable": summary.promotable,
})
.to_string()
}
#[wasm_bindgen(js_name = optimizeJson)]
pub fn optimize_json(&self, campaign_json: &str) -> 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}")),
};
match self.inner.optimize(&spec) {
Ok(report) => {
serde_json::to_string(&report).unwrap_or_else(|e| embedder_json(&e.to_string()))
}
Err(e) => error_json(&e),
}
}
#[wasm_bindgen(js_name = exportBankJson)]
pub fn export_bank_json(&self) -> String {
match self.inner.export_bank() {
Ok(s) => s,
Err(e) => error_json(&e),
}
}
#[wasm_bindgen(js_name = importBankJson)]
pub fn import_bank_json(&mut self, bank_json: &str) -> String {
match self.inner.import_bank(bank_json) {
Ok(()) => "{\"ok\":true}".to_string(),
Err(e) => error_json(&e),
}
}
}
#[cfg(feature = "wasm-onnx")]
#[derive(Deserialize)]
struct FromBytesOptions {
manifest: serde_json::Value,
}
#[cfg(feature = "wasm-onnx")]
fn build_onnx(
options_json: &str,
model_bytes: &[u8],
tokenizer_bytes: &[u8],
) -> std::result::Result<Engine, JsValue> {
use ruvector_embed_core::{ModelManifest, TractEmbedder};
let opts: FromBytesOptions = serde_json::from_str(options_json)
.map_err(|e| JsValue::from_str(&format!("invalid options JSON: {e}")))?;
let manifest_json = match opts.manifest {
serde_json::Value::String(s) => s,
other => other.to_string(),
};
let manifest = ModelManifest::from_json(&manifest_json)
.map_err(|e| JsValue::from_str(&format!("manifest: {e}")))?;
let embedder = TractEmbedder::from_bytes(model_bytes, tokenizer_bytes, &manifest)
.map_err(|e| JsValue::from_str(&format!("onnx load: {e}")))?;
Ok(Engine {
inner: CoreEngine::new(Box::new(embedder)),
})
}
#[cfg(not(feature = "wasm-onnx"))]
fn build_onnx(
_options_json: &str,
_model_bytes: &[u8],
_tokenizer_bytes: &[u8],
) -> std::result::Result<Engine, JsValue> {
Err(JsValue::from_str(&onnx_error_json()))
}
#[cfg(not(feature = "wasm-onnx"))]
fn onnx_error_json() -> String {
embedder_json(
"onnx backend requires ruvector-embed-core (tract, wasm) — rebuild with --features wasm-onnx",
)
}
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()
}