1pub mod declaration;
32
33use std::sync::Arc;
34
35use wasmtime::{FuncType, HeapType, Linker, RefType, StructType, Val, ValType};
36
37use crate::runtime::StoreData;
38use crate::runtime::call_log::{ModelUsage, Payload, Side, record_payload, record_usage};
39use crate::runtime::decision::CallTicket;
40use crate::runtime::embedding::{
41 EmbeddingBatch, EmbeddingBoundKind, EmbeddingError, EmbeddingLimits, EmbeddingMalformedReason,
42 EmbeddingModel, EmbeddingProvider, EmbeddingTokenBudget, Purpose,
43};
44use crate::runtime::fuel;
45use crate::runtime::host::{
46 abi_arg, abi_result, fatal_host_error, quota_exceeded_error, range_error, read_string_arg,
47 register_host_fn, register_host_fn_async, type_error, write_boxed_number_struct,
48 write_submilli_array_struct_precharged, write_submilli_string_struct,
49 write_submilli_uint8array_struct, write_uint8_array_precharged,
50};
51use crate::runtime::intrinsic_types::build_intrinsic_types;
52use crate::stdlib::abi::{
53 self, backing_struct, f64_field, install_field_getters, nullable_object_field, raw_bytes_field,
54 string_field,
55};
56use crate::stdlib::shared::{
57 audit_quota_denial, check_security_call, filters_candidate, mark_filtered, optional_number,
58 preflight_models, sanitize_description,
59};
60
61pub use crate::runtime::EMBEDDING_MODULE_NAME as MODULE_NAME;
62pub use declaration::package_declaration;
63
64pub const CAPABILITY: &str = "embedding.embed";
66
67const E_VECTORS: usize = 1;
69const E_COUNT: usize = 2;
70const E_DIMENSIONS: usize = 3;
71const E_IDENTITY: usize = 4;
72const E_MODEL: usize = 5;
73const E_INPUT_TOKENS: usize = 6;
74
75const M_NAME: usize = 1;
77const M_DESCRIPTION: usize = 2;
78const M_DIMENSIONS: usize = 3;
79const M_MAX_INPUT_TOKENS: usize = 4;
80const M_MAX_INPUT_BYTES: usize = 5;
81const M_IDENTITY: usize = 6;
82
83pub fn install(linker: &mut Linker<StoreData>) -> wasmtime::Result<()> {
84 let engine = linker.engine().clone();
85 let intr = build_intrinsic_types(&engine)?;
86 let string = ValType::Ref(RefType::new(
87 false,
88 HeapType::ConcreteStruct(intr.string.clone()),
89 ));
90 let array = ValType::Ref(RefType::new(
91 false,
92 HeapType::ConcreteStruct(intr.array.clone()),
93 ));
94 let uint8 = ValType::Ref(RefType::new(
95 false,
96 HeapType::ConcreteStruct(intr.uint8_array.clone()),
97 ));
98 let nullable_object = ValType::Ref(RefType::new(
101 true,
102 HeapType::ConcreteStruct(intr.object.clone()),
103 ));
104 let receiver = ValType::Ref(RefType::new(
105 false,
106 HeapType::ConcreteStruct(intr.object.clone()),
107 ));
108
109 register_host_fn_async(
110 linker,
111 MODULE_NAME,
112 crate::mangle::package_symbol(MODULE_NAME, "embed"),
113 FuncType::new(
114 &engine,
115 [string.clone(), array.clone(), string.clone()],
116 [nullable_object.clone()],
117 ),
118 false,
119 |caller, params, results| {
120 Box::pin(async move {
121 let model =
122 read_string_arg(&mut *caller, abi_arg(params, 0)?, "embedding.embed (model)")?;
123 let purpose = read_string_arg(
124 &mut *caller,
125 abi_arg(params, 2)?,
126 "embedding.embed (purpose)",
127 )?;
128 *abi_result(results, 0)? =
129 embed(caller, &model, abi_arg(params, 1)?, &purpose).await?;
130 Ok(())
131 })
132 },
133 )?;
134
135 register_host_fn_async(
136 linker,
137 MODULE_NAME,
138 crate::mangle::package_symbol(MODULE_NAME, "models"),
139 FuncType::new(&engine, [], [array.clone()]),
140 false,
141 |caller, _params, results| {
142 Box::pin(async move {
143 *abi_result(results, 0)? = models(caller).await?;
144 Ok(())
145 })
146 },
147 )?;
148
149 install_getters(linker, &engine, &receiver, string, nullable_object)?;
150 install_methods(linker, &engine, receiver, array, uint8)
151}
152
153fn install_getters(
154 linker: &mut Linker<StoreData>,
155 engine: &wasmtime::Engine,
156 receiver: &ValType,
157 string: ValType,
158 nullable_object: ValType,
159) -> wasmtime::Result<()> {
160 install_field_getters(
161 linker,
162 MODULE_NAME,
163 "Embeddings",
164 engine,
165 receiver,
166 &[
167 ("count", E_COUNT, ValType::F64),
168 ("dimensions", E_DIMENSIONS, ValType::F64),
169 ("identity", E_IDENTITY, string.clone()),
170 ("model", E_MODEL, string.clone()),
171 ("inputTokens", E_INPUT_TOKENS, nullable_object.clone()),
172 ],
173 )?;
174 install_field_getters(
175 linker,
176 MODULE_NAME,
177 "EmbeddingModel",
178 engine,
179 receiver,
180 &[
181 ("name", M_NAME, string.clone()),
182 ("description", M_DESCRIPTION, nullable_object.clone()),
183 ("dimensions", M_DIMENSIONS, ValType::F64),
184 ("maxInputTokens", M_MAX_INPUT_TOKENS, nullable_object),
185 ("maxInputBytes", M_MAX_INPUT_BYTES, ValType::F64),
186 ("identity", M_IDENTITY, string),
187 ],
188 )
189}
190
191fn install_methods(
192 linker: &mut Linker<StoreData>,
193 engine: &wasmtime::Engine,
194 receiver: ValType,
195 array: ValType,
196 uint8: ValType,
197) -> wasmtime::Result<()> {
198 let embeddings_key = crate::mangle::package_symbol(MODULE_NAME, "Embeddings");
199
200 register_host_fn(
201 linker,
202 MODULE_NAME,
203 crate::mangle::extend(&embeddings_key, "vector"),
204 FuncType::new(engine, [receiver.clone(), ValType::F64], [array]),
205 true,
206 |caller, params, results| {
207 let row = read_row(
208 caller,
209 abi_arg(params, 0)?,
210 abi_arg(params, 1)?,
211 "embedding.vector",
212 )?;
213 *abi_result(results, 0)? = build_number_array(caller, &row)?;
214 Ok(())
215 },
216 )?;
217
218 register_host_fn(
219 linker,
220 MODULE_NAME,
221 crate::mangle::extend(&embeddings_key, "bytes"),
222 FuncType::new(engine, [receiver, ValType::F64], [uint8]),
223 true,
224 |caller, params, results| {
225 let row = read_row(
226 caller,
227 abi_arg(params, 0)?,
228 abi_arg(params, 1)?,
229 "embedding.bytes",
230 )?;
231 *abi_result(results, 0)? = Val::AnyRef(Some(
232 write_submilli_uint8array_struct(caller, &row)?.to_anyref(),
233 ));
234 Ok(())
235 },
236 )
237}
238
239async fn embed(
253 caller: &mut wasmtime::Caller<'_, StoreData>,
254 model: &str,
255 texts: &Val,
256 purpose: &str,
257) -> wasmtime::Result<Val> {
258 let elements = crate::runtime::prelude::collection::read_array_vals(caller, texts)?;
259 let ticket = gate(caller, model, elements.len())?;
260
261 let budget = execution_budget(caller);
262 let limits = budget.limits();
263 let texts = read_texts(caller, model, &elements, &limits)?;
264 record_payload(&*caller, ticket, Side::Request, || {
267 let body = serde_json::to_vec(&texts).unwrap_or_default();
268 let bytes: usize = texts.iter().map(String::len).sum();
269 Payload::meta(serde_json::json!({
270 "op": "embed",
271 "model": model,
272 "purpose": purpose,
273 "count": texts.len(),
274 }))
275 .with_owned_body(body)
276 .with_size(bytes as u64)
277 });
278 let purpose = parse_purpose(model, purpose)?;
279
280 let provider = provider(caller, model)?;
281 check_inputs(caller, &*provider, model, &texts).await?;
282
283 let sent: usize = texts.iter().map(String::len).sum();
284 fuel::charge(&mut *caller, fuel::IO, sent as u64)?;
285
286 let estimate = provider.estimate_tokens(model, &texts);
287 budget
288 .reserve(model, estimate)
289 .map_err(|error| quota_throw(caller, ticket, model, error))?;
290
291 let batch = match provider.embed(model, &texts, purpose, &budget).await {
292 Ok(batch) => batch,
293 Err(error) => {
294 budget.settle(estimate, error.settlements());
297 let reported: u64 = error
302 .settlements()
303 .iter()
304 .fold(0, |sum, settlement| sum.saturating_add(settlement.reported));
305 if reported > 0 {
306 record_usage(
307 &*caller,
308 ticket,
309 ModelUsage {
310 input_tokens: Some(reported),
311 output_tokens: None,
312 },
313 );
314 }
315 return Err(quota_throw(caller, ticket, model, error));
316 }
317 };
318 budget.settle(estimate, batch.settlements());
319
320 if batch.count() != texts.len() {
323 return Err(throw(EmbeddingError::Malformed {
324 alias: model.to_string(),
325 reason: EmbeddingMalformedReason::CountMismatch,
326 settlements: Vec::new(),
327 }));
328 }
329
330 let received = (batch.values().len() as u64).saturating_mul(4);
331 record_payload(&*caller, ticket, Side::Response, || {
333 Payload::meta(serde_json::json!({
334 "count": batch.count(),
335 "dimensions": batch.dimensions(),
336 "identity": batch.identity(),
337 "model": model,
338 "inputTokens": batch.input_tokens(),
339 }))
340 .with_size(received)
341 });
342 record_usage(
343 &*caller,
344 ticket,
345 ModelUsage {
346 input_tokens: batch.input_tokens(),
347 output_tokens: None,
348 },
349 );
350 fuel::settle(&mut *caller, fuel::IO, received)?;
351 build_embeddings(caller, batch)
352}
353
354fn gate(
357 caller: &mut wasmtime::Caller<'_, StoreData>,
358 model: &str,
359 input_count: usize,
360) -> wasmtime::Result<Option<CallTicket>> {
361 check_security_call(
362 caller,
363 CAPABILITY,
364 serde_json::json!({ "model": model, "input_count": input_count }),
365 )
366}
367
368fn read_texts(
371 caller: &mut wasmtime::Caller<'_, StoreData>,
372 model: &str,
373 elements: &[Val],
374 limits: &EmbeddingLimits,
375) -> wasmtime::Result<Vec<String>> {
376 let count = elements.len() as u64;
377 if count > limits.max_texts_per_call {
378 return Err(throw(EmbeddingError::BoundsExceeded {
379 alias: model.to_string(),
380 kind: EmbeddingBoundKind::TextCount,
381 actual: count,
382 limit: limits.max_texts_per_call,
383 }));
384 }
385 if elements.is_empty() {
386 return Err(range_error(format!(
387 "embedding.embed(\"{model}\"): at least one text is required — pass a non-empty array"
388 )));
389 }
390 let mut texts = Vec::with_capacity(elements.len());
391 let mut total = 0u64;
392 for element in elements {
393 let text = read_string_arg(caller, element, "embedding.embed (texts)")?;
394 total = total.saturating_add(text.len() as u64);
395 if total > limits.max_bytes_per_call {
396 return Err(throw(EmbeddingError::BoundsExceeded {
397 alias: model.to_string(),
398 kind: EmbeddingBoundKind::TotalBytes,
399 actual: total,
400 limit: limits.max_bytes_per_call,
401 }));
402 }
403 texts.push(text);
404 }
405 Ok(texts)
406}
407
408fn parse_purpose(model: &str, purpose: &str) -> wasmtime::Result<Purpose> {
409 match purpose {
410 "query" => Ok(Purpose::Query),
411 "document" => Ok(Purpose::Document),
412 _ => Err(range_error(format!(
413 "embedding.embed(\"{model}\"): purpose must be \"query\" or \"document\""
414 ))),
415 }
416}
417
418async fn check_inputs(
421 caller: &mut wasmtime::Caller<'_, StoreData>,
422 provider: &dyn EmbeddingProvider,
423 model: &str,
424 texts: &[String],
425) -> wasmtime::Result<()> {
426 let Some(limit) = provider.max_input_bytes(model) else {
427 let available = available_aliases(caller, provider).await?;
428 return Err(throw(EmbeddingError::UnknownModel {
429 alias: model.to_string(),
430 available,
431 }));
432 };
433 let too_long = texts.iter().position(|text| text.len() as u64 > limit);
437 if let Some(index) = too_long {
438 return Err(throw(EmbeddingError::InputTooLong {
439 alias: model.to_string(),
440 index: Some(index),
441 limit: Some(limit),
442 settlements: Vec::new(),
443 }));
444 }
445 Ok(())
446}
447
448async fn available_aliases(
452 caller: &mut wasmtime::Caller<'_, StoreData>,
453 provider: &dyn EmbeddingProvider,
454) -> wasmtime::Result<Vec<String>> {
455 let candidates = provider.models().await.map_err(throw)?;
456 let mut names = Vec::new();
457 for candidate in candidates {
458 if may_embed(caller, &candidate.name)? {
459 names.push(candidate.name);
460 }
461 }
462 Ok(names)
463}
464
465async fn models(caller: &mut wasmtime::Caller<'_, StoreData>) -> wasmtime::Result<Val> {
472 preflight_models(caller, CAPABILITY, "input_count")?;
473 let provider = provider(caller, "")?;
474 let candidates = provider.models().await.map_err(throw)?;
475
476 let mut built = Vec::with_capacity(candidates.len());
477 for candidate in candidates {
478 if may_embed(caller, &candidate.name)? {
479 built.push(build_model(caller, candidate)?);
480 }
481 }
482 abi::new_array(caller, &built)
483}
484
485fn may_embed(caller: &mut wasmtime::Caller<'_, StoreData>, model: &str) -> wasmtime::Result<bool> {
489 let keeps = filters_candidate(gate(caller, model, 0).map(|_| ()))?;
490 if !keeps {
491 mark_filtered(&*caller);
492 }
493 Ok(keeps)
494}
495
496fn provider(
499 caller: &wasmtime::Caller<'_, StoreData>,
500 model: &str,
501) -> wasmtime::Result<Arc<dyn EmbeddingProvider>> {
502 caller.data().embedding_provider.clone().ok_or_else(|| {
503 throw(EmbeddingError::NotConfigured {
504 alias: model.to_string(),
505 })
506 })
507}
508
509fn execution_budget(caller: &wasmtime::Caller<'_, StoreData>) -> Arc<EmbeddingTokenBudget> {
512 caller
513 .data()
514 .embedding_budget
515 .clone()
516 .unwrap_or_else(|| Arc::new(EmbeddingTokenBudget::unmetered()))
517}
518
519fn quota_throw(
521 caller: &wasmtime::Caller<'_, StoreData>,
522 ticket: Option<CallTicket>,
523 model: &str,
524 error: EmbeddingError,
525) -> wasmtime::Error {
526 if !error.is_budget_exceeded() {
527 return throw(error);
528 }
529 if let Err(denial) = audit_quota_denial(
530 caller,
531 ticket,
532 CAPABILITY,
533 model,
534 "embedding-token budget exceeded",
535 ) {
536 return denial;
537 }
538 throw(error)
539}
540
541fn throw(error: EmbeddingError) -> wasmtime::Error {
550 let message = error.to_string();
551 match error {
552 EmbeddingError::BudgetExceeded { .. } => quota_exceeded_error(message),
553 EmbeddingError::BoundsExceeded { .. } | EmbeddingError::InputTooLong { .. } => {
554 range_error(message)
555 }
556 EmbeddingError::Malformed { .. } => type_error(message),
557 EmbeddingError::Internal { .. } => fatal_host_error(message),
558 EmbeddingError::NotConfigured { .. }
559 | EmbeddingError::UnknownModel { .. }
560 | EmbeddingError::Unauthorized { .. }
561 | EmbeddingError::Provider { .. } => wasmtime::Error::msg(message),
562 }
563}
564
565fn embeddings_backing_struct(engine: &wasmtime::Engine) -> wasmtime::Result<StructType> {
566 let intr = build_intrinsic_types(engine)?;
567 backing_struct(
568 engine,
569 &intr,
570 vec![
571 raw_bytes_field(&intr), f64_field(), f64_field(), string_field(&intr), string_field(&intr), nullable_object_field(&intr), ],
578 )
579}
580
581fn model_backing_struct(engine: &wasmtime::Engine) -> wasmtime::Result<StructType> {
582 let intr = build_intrinsic_types(engine)?;
583 backing_struct(
584 engine,
585 &intr,
586 vec![
587 string_field(&intr), nullable_object_field(&intr), f64_field(), nullable_object_field(&intr), f64_field(), string_field(&intr), ],
594 )
595}
596
597fn build_embeddings(
606 caller: &mut wasmtime::Caller<'_, StoreData>,
607 batch: EmbeddingBatch,
608) -> wasmtime::Result<Val> {
609 let byte_count = batch
610 .values()
611 .len()
612 .checked_mul(4)
613 .ok_or_else(|| fatal_host_error("embedding.embed: result size overflows"))?;
614 fuel::settle(&mut *caller, fuel::COPY, byte_count as u64)?;
615 let mut bytes = Vec::new();
616 bytes
617 .try_reserve_exact(byte_count)
618 .map_err(|error| fatal_host_error(format!("embedding.embed: {error}")))?;
619 for value in batch.values() {
620 bytes.extend_from_slice(&value.to_le_bytes());
621 }
622 let count = batch.count() as f64;
623 let dimensions = batch.dimensions() as f64;
624 let identity = batch.identity().to_string();
625 let model = batch.model().to_string();
626 let input_tokens = batch.input_tokens().map(|n| n as f64);
627 drop(batch);
628
629 let array_ty = build_intrinsic_types(caller.engine())?.raw_uint8_array;
631 let vectors = write_uint8_array_precharged(&mut *caller, array_ty, &bytes)?;
632 drop(bytes);
633
634 let identity = string_val(caller, &identity)?;
635 let model = string_val(caller, &model)?;
636 let input_tokens = optional_number(caller, input_tokens)?;
637
638 let ty = embeddings_backing_struct(caller.engine())?;
639 abi::new_backing(
640 caller,
641 ty,
642 &[
643 Val::AnyRef(Some(vectors.to_anyref())),
644 Val::F64(count.to_bits()),
645 Val::F64(dimensions.to_bits()),
646 identity,
647 model,
648 input_tokens,
649 ],
650 )
651}
652
653fn build_model(
654 caller: &mut wasmtime::Caller<'_, StoreData>,
655 model: EmbeddingModel,
656) -> wasmtime::Result<Val> {
657 let name = string_val(caller, &model.name)?;
658 let description = match model.description.as_deref().and_then(sanitize_description) {
659 Some(text) => string_val(caller, &text)?,
660 None => crate::runtime::prelude::undefined::value(caller)?,
661 };
662 let max_input_tokens = optional_number(caller, model.max_input_tokens.map(|n| n as f64))?;
663 let identity = string_val(caller, &model.identity)?;
664 let ty = model_backing_struct(caller.engine())?;
665 abi::new_backing(
666 caller,
667 ty,
668 &[
669 name,
670 description,
671 Val::F64((model.dimensions as f64).to_bits()),
672 max_input_tokens,
673 Val::F64((model.max_input_bytes as f64).to_bits()),
674 identity,
675 ],
676 )
677}
678
679fn string_val(caller: &mut wasmtime::Caller<'_, StoreData>, text: &str) -> wasmtime::Result<Val> {
680 Ok(Val::AnyRef(Some(
681 write_submilli_string_struct(caller, text)?.to_anyref(),
682 )))
683}
684
685fn read_row(
688 caller: &mut wasmtime::Caller<'_, StoreData>,
689 receiver: &Val,
690 index: &Val,
691 ctx: &str,
692) -> wasmtime::Result<Vec<u8>> {
693 let Val::F64(bits) = index else {
694 return Err(fatal_host_error(format!("{ctx}: index is not a number")));
695 };
696 let index = f64::from_bits(*bits);
697
698 let st = abi::backing_receiver(caller, receiver)?;
699 let Val::AnyRef(Some(vectors)) = st.field(&mut *caller, E_VECTORS)? else {
700 return Err(fatal_host_error(format!(
701 "{ctx}: result vectors are missing"
702 )));
703 };
704 let vectors = vectors
705 .as_array(&mut *caller)?
706 .ok_or_else(|| fatal_host_error(format!("{ctx}: result vectors are not an array")))?;
707 let (Val::F64(count), Val::F64(dimensions)) = (
708 st.field(&mut *caller, E_COUNT)?,
709 st.field(&mut *caller, E_DIMENSIONS)?,
710 ) else {
711 return Err(fatal_host_error(format!("{ctx}: result shape is missing")));
712 };
713 let (count, dimensions) = (f64::from_bits(count), f64::from_bits(dimensions));
714 let in_range = index.is_finite() && index.fract() == 0.0 && index >= 0.0 && index < count;
715 if !in_range {
716 return Err(range_error(format!(
717 "{ctx}: index {index} is out of range for {count} vectors — use an integer from 0 to \
718 count - 1"
719 )));
720 }
721 let row_bytes = (dimensions as usize)
722 .checked_mul(4)
723 .ok_or_else(|| fatal_host_error(format!("{ctx}: row size overflows")))?;
724 let offset = (index as usize)
725 .checked_mul(row_bytes)
726 .and_then(|offset| u32::try_from(offset).ok())
727 .ok_or_else(|| fatal_host_error(format!("{ctx}: row offset overflows")))?;
728 fuel::charge(&mut *caller, fuel::COPY, row_bytes as u64)?;
730 let mut row = Vec::new();
731 row.try_reserve_exact(row_bytes)
732 .map_err(|error| fatal_host_error(format!("{ctx}: {error}")))?;
733 row.resize(row_bytes, 0);
734 vectors
735 .read_i8(&mut *caller, offset, &mut row)
736 .map_err(|error| fatal_host_error(format!("{ctx}: {error}")))?;
737 Ok(row)
738}
739
740fn build_number_array(
742 caller: &mut wasmtime::Caller<'_, StoreData>,
743 row: &[u8],
744) -> wasmtime::Result<Val> {
745 fuel::charge(&mut *caller, fuel::ELEM, (row.len() / 4) as u64)?;
747 let mut boxed = Vec::with_capacity(row.len() / 4);
748 for chunk in row.as_chunks::<4>().0 {
749 let value = f32::from_le_bytes(*chunk);
750 boxed.push(Val::AnyRef(Some(
751 write_boxed_number_struct(caller, f64::from(value))?.to_anyref(),
752 )));
753 }
754 Ok(Val::AnyRef(Some(
755 write_submilli_array_struct_precharged(caller, &boxed)?.to_anyref(),
756 )))
757}
758
759#[cfg(test)]
760mod tests {
761 use std::pin::Pin;
762 use std::sync::atomic::{AtomicUsize, Ordering};
763 use std::sync::{Arc, Mutex};
764
765 use crate::runtime::{
766 CheckOutcome, EmbeddingBatch, EmbeddingError, EmbeddingFailureReason, EmbeddingLimits,
767 EmbeddingModel, EmbeddingProvider, EmbeddingTokenBudget, Purpose, RuntimeConfig,
768 SecurityCheck, SharedTokenBudget, StoreData, SubBatchSettlement, Vfs, dispatch_main_async,
769 install_runtime_async,
770 };
771
772 type BoxFuture<'a, T> = Pin<Box<dyn std::future::Future<Output = T> + Send + 'a>>;
773
774 #[derive(Clone, Copy)]
776 enum Outcome {
777 Reported,
779 TransportFailure,
781 Internal,
783 }
784
785 struct MockProvider {
788 dimensions: usize,
789 max_input_bytes: u64,
790 outcome: Outcome,
791 embed_calls: AtomicUsize,
792 }
793
794 impl MockProvider {
795 fn new(dimensions: usize, max_input_bytes: u64, outcome: Outcome) -> Arc<Self> {
796 Arc::new(Self {
797 dimensions,
798 max_input_bytes,
799 outcome,
800 embed_calls: AtomicUsize::new(0),
801 })
802 }
803
804 fn calls(&self) -> usize {
805 self.embed_calls.load(Ordering::Relaxed)
806 }
807 }
808
809 impl EmbeddingProvider for MockProvider {
810 fn embed<'a>(
811 &'a self,
812 alias: &'a str,
813 texts: &'a [String],
814 _purpose: Purpose,
815 budget: &'a EmbeddingTokenBudget,
816 ) -> BoxFuture<'a, Result<EmbeddingBatch, EmbeddingError>> {
817 self.embed_calls.fetch_add(1, Ordering::Relaxed);
818 Box::pin(async move {
819 let estimate = self.estimate_tokens(alias, texts);
820 budget.mark_sent(alias, estimate)?;
821 if matches!(self.outcome, Outcome::Internal) {
822 return Err(EmbeddingError::Internal {
823 alias: alias.to_string(),
824 settlements: vec![SubBatchSettlement {
825 estimate,
826 reported: 0,
827 indeterminate: estimate,
828 }],
829 });
830 }
831 if matches!(self.outcome, Outcome::TransportFailure) {
832 return Err(EmbeddingError::Provider {
833 alias: alias.to_string(),
834 reason: EmbeddingFailureReason::Transport,
835 settlements: vec![SubBatchSettlement {
836 estimate,
837 reported: 0,
838 indeterminate: estimate,
839 }],
840 });
841 }
842 let values = vec![0.5f32; texts.len() * self.dimensions];
843 EmbeddingBatch::new(values, texts.len(), self.dimensions, "emb1:mock", alias)
844 .map(|batch| {
845 batch.with_settlements(vec![SubBatchSettlement {
846 estimate,
847 reported: estimate,
848 indeterminate: 0,
849 }])
850 })
851 .map_err(|_| EmbeddingError::Malformed {
852 alias: alias.to_string(),
853 reason: crate::runtime::EmbeddingMalformedReason::InvalidBody,
854 settlements: Vec::new(),
855 })
856 })
857 }
858
859 fn models<'a>(&'a self) -> BoxFuture<'a, Result<Vec<EmbeddingModel>, EmbeddingError>> {
860 let model = EmbeddingModel {
861 name: "mock".to_string(),
862 description: None,
863 dimensions: self.dimensions as u64,
864 max_input_tokens: None,
865 max_input_bytes: self.max_input_bytes,
866 identity: "emb1:mock".to_string(),
867 };
868 Box::pin(async move { Ok(vec![model]) })
869 }
870
871 fn max_input_bytes(&self, alias: &str) -> Option<u64> {
872 (alias == "mock").then_some(self.max_input_bytes)
873 }
874 }
875
876 struct RecordingPolicy(Arc<Mutex<Vec<serde_json::Value>>>);
877
878 impl SecurityCheck for RecordingPolicy {
879 fn check(
880 &self,
881 _caller: &str,
882 _capability: &str,
883 context: &serde_json::Value,
884 ) -> CheckOutcome {
885 if let Ok(mut contexts) = self.0.lock() {
886 contexts.push(context.clone());
887 }
888 CheckOutcome::Allow { rule: None }
889 }
890 }
891
892 struct Run {
895 result: wasmtime::Result<String>,
896 host_attached_bytes: u64,
897 observed_bytes: u64,
899 }
900
901 async fn run(
902 source: &str,
903 provider: Option<Arc<MockProvider>>,
904 budget: Option<Arc<EmbeddingTokenBudget>>,
905 policy: Option<Arc<dyn SecurityCheck>>,
906 ) -> Run {
907 let compiled = crate::compile_script(source, "test.ts", crate::FileId(0), &[], &[])
908 .expect("compile clean");
909 let cfg = RuntimeConfig::default();
910 let engine = cfg.engine().expect("engine");
911 let mut data = StoreData::with_vfs(Vfs::tempdir().expect("tempdir"));
912 data.install_type_info(compiled.type_info.clone());
913 data.embedding_provider = provider.map(|p| p as Arc<dyn EmbeddingProvider>);
914 data.embedding_budget = budget;
915 if let Some(policy) = policy {
916 data.security_check = policy;
917 }
918 let mut store = cfg.store_async(&engine, data).expect("store");
919 crate::runtime::install_tenant_limits(&mut store);
920 let module = wasmtime::Module::new(&engine, &compiled.wasm).expect("module");
921 let mut linker = wasmtime::Linker::<StoreData>::new(&engine);
922 install_runtime_async(&mut linker, &mut store)
923 .await
924 .expect("install");
925 let inst = linker
926 .instantiate_async(&mut store, &module)
927 .await
928 .expect("instantiate");
929 let result = dispatch_main_async(&mut store, &inst)
930 .await
931 .map(Option::unwrap_or_default);
932 let host_attached_bytes = store.data().tenant_limits.host_attached_bytes();
933 let observed_bytes = store.data().tenant_limits.observed_bytes();
934 Run {
935 result,
936 host_attached_bytes,
937 observed_bytes,
938 }
939 }
940
941 fn budget(limits: EmbeddingLimits) -> Arc<EmbeddingTokenBudget> {
942 Arc::new(EmbeddingTokenBudget::new(
943 limits,
944 SharedTokenBudget::new(u64::MAX),
945 ))
946 }
947
948 const ROUND_TRIP: &str = r#"import embedding from "submilli:embedding";
949 function main(): void {
950 const r = embedding.embed("mock", ["a", "b", "c"], "document");
951 assert(r.count === 3, "three vectors");
952 }"#;
953
954 #[tokio::test]
957 async fn ae1_over_length_input_is_refused_before_the_provider_and_the_budget() {
958 let provider = MockProvider::new(4, 10, Outcome::Reported);
959 let budget = budget(EmbeddingLimits::default());
960 let run = run(
961 r#"import embedding from "submilli:embedding";
962 function main(): void {
963 let message = "";
964 try {
965 embedding.embed("mock", ["ok", "ok", "this one is too long"], "document");
966 } catch (e: RangeError) {
967 message = e.message;
968 }
969 assert(message.indexOf("text 2") >= 0, "names index 2: " + message);
970 }"#,
971 Some(Arc::clone(&provider)),
972 Some(Arc::clone(&budget)),
973 None,
974 )
975 .await;
976 run.result.expect("program completes");
977 assert_eq!(provider.calls(), 0, "the provider is never called");
978 assert_eq!(budget.used(), 0, "nothing is charged");
979 assert_eq!(budget.held(), 0, "nothing is held");
980 assert_eq!(budget.requests(), 0, "nothing counts as sent");
981 }
982
983 #[tokio::test]
986 async fn ae6_exhausted_budget_refuses_without_calling_the_provider() {
987 let provider = MockProvider::new(4, 1_000, Outcome::Reported);
988 let budget = budget(EmbeddingLimits {
989 per_execution_tokens: 1,
990 ..EmbeddingLimits::default()
991 });
992 let run = run(
993 r#"import embedding from "submilli:embedding";
994 function main(): void {
995 let quota = false;
996 try {
997 embedding.embed("mock", ["more than one token of text"], "document");
998 } catch (e: QuotaExceededError) {
999 quota = true;
1000 }
1001 assert(quota, "the call is a QuotaExceededError");
1002 }"#,
1003 Some(Arc::clone(&provider)),
1004 Some(Arc::clone(&budget)),
1005 None,
1006 )
1007 .await;
1008 run.result.expect("program completes");
1009 assert_eq!(provider.calls(), 0, "the provider is never called");
1010 assert_eq!(budget.used(), 0, "the refusal charged nothing");
1011 assert_eq!(budget.requests(), 0);
1012 }
1013
1014 #[tokio::test]
1017 async fn success_settles_reported_usage() {
1018 let provider = MockProvider::new(4, 1_000, Outcome::Reported);
1019 let budget = budget(EmbeddingLimits::default());
1020 let run = run(ROUND_TRIP, Some(provider), Some(Arc::clone(&budget)), None).await;
1021 run.result.expect("program completes");
1022 assert_eq!(
1023 budget.used(),
1024 3,
1025 "three one-byte texts estimate one token each"
1026 );
1027 assert_eq!(budget.held(), 0, "reported usage leaves nothing held");
1028 assert_eq!(budget.requests(), 1);
1029 }
1030
1031 #[tokio::test]
1034 async fn failure_after_send_settles_the_estimate_as_held() {
1035 let provider = MockProvider::new(4, 1_000, Outcome::TransportFailure);
1036 let budget = budget(EmbeddingLimits::default());
1037 let run = run(
1038 r#"import embedding from "submilli:embedding";
1039 function main(): void {
1040 let message = "";
1041 try {
1042 embedding.embed("mock", ["secret text"], "document");
1043 } catch (e: Error) {
1044 message = e.message;
1045 }
1046 assert(message.indexOf("transport") >= 0, message);
1047 assert(message.indexOf("secret text") < 0, "no input text in the error");
1048 }"#,
1049 Some(provider),
1050 Some(Arc::clone(&budget)),
1051 None,
1052 )
1053 .await;
1054 run.result.expect("program completes");
1055 assert_eq!(budget.held(), 4, "ceil(11 / 3) tokens stay held");
1056 assert_eq!(budget.used(), 4, "and count against the run");
1057 }
1058
1059 #[tokio::test]
1063 async fn ae8_results_are_held_at_four_bytes_per_number() {
1064 const COUNT: u64 = 128;
1065 const DIMENSIONS: u64 = 3072;
1066 let source = |texts: u32| {
1067 format!(
1068 r#"import embedding from "submilli:embedding";
1069 function main(): void {{
1070 const texts: string[] = [];
1071 for (let i = 0; i < {texts}; i = i + 1) {{
1072 texts.push("t");
1073 }}
1074 const r = embedding.embed("mock", texts, "document");
1075 assert(r.count === {texts} && r.dimensions === 3072, "shape");
1076 assert(r.vector(0).length === 3072, "reading vector 0 returns 3,072 numbers");
1077 assert(r.bytes(0).length === 12288, "exporting it returns 12,288 bytes");
1078 if ({texts} > 5) {{
1079 assert(r.vector(5).length === 3072, "reading vector 5 returns 3,072 numbers");
1080 assert(r.bytes(5).length === 12288, "exporting it returns 12,288 bytes");
1081 }}
1082 let range = false;
1083 try {{
1084 r.vector({texts});
1085 }} catch (e: RangeError) {{
1086 range = true;
1087 }}
1088 assert(range, "reading vector {texts} is a RangeError");
1089 }}"#
1090 )
1091 };
1092 let provider = |_| MockProvider::new(DIMENSIONS as usize, 1_000, Outcome::Reported);
1093 let one = run(&source(1), Some(provider(())), None, None).await;
1094 one.result.expect("one-vector program completes");
1095 let many = run(&source(COUNT as u32), Some(provider(())), None, None).await;
1096 many.result.expect("program completes");
1097
1098 let expected = (COUNT - 1) * DIMENSIONS * 4;
1102 let grown = many.observed_bytes.saturating_sub(one.observed_bytes);
1103 assert!(
1104 grown >= expected - expected / 10 && grown <= expected + expected / 10,
1105 "the result costs about count x dimensions x 4 bytes of GC heap: grew {grown}, \
1106 expected about {expected}"
1107 );
1108 assert_eq!(
1109 many.host_attached_bytes, 0,
1110 "the vectors are not host-attached bytes"
1111 );
1112 }
1113
1114 #[tokio::test]
1119 async fn discarded_results_are_collected_under_the_default_cap() {
1120 let provider = MockProvider::new(3072, 1_000, Outcome::Reported);
1121 let run = run(
1122 r#"import embedding from "submilli:embedding";
1123 function main(): void {
1124 const texts: string[] = [];
1125 for (let i = 0; i < 128; i = i + 1) {
1126 texts.push("t");
1127 }
1128 let total = 0;
1129 for (let round = 0; round < 100; round = round + 1) {
1130 const r = embedding.embed("mock", texts, "document");
1131 total = total + r.count;
1132 }
1133 assert(total === 12800, "every round returned its vectors");
1134 }"#,
1135 Some(provider),
1136 Some(budget(EmbeddingLimits::default())),
1137 None,
1138 )
1139 .await;
1140 run.result
1141 .expect("program completes without memory exhaustion");
1142 }
1143
1144 #[tokio::test]
1147 async fn the_hidden_vectors_are_not_reachable_from_the_guest() {
1148 for member in ["vectors", "handle", "data", "buffer"] {
1149 let source = format!(
1150 r#"import embedding from "submilli:embedding";
1151 function main(): void {{
1152 const r = embedding.embed("mock", ["a"], "document");
1153 const hidden = r.{member};
1154 }}"#
1155 );
1156 assert!(
1157 crate::compile_script(&source, "test.ts", crate::FileId(0), &[], &[]).is_err(),
1158 "`Embeddings.{member}` must not type-check"
1159 );
1160 }
1161 let provider = MockProvider::new(4, 1_000, Outcome::Reported);
1162 let cast = run(
1163 r#"import embedding from "submilli:embedding";
1164 function main(): void {
1165 const r = embedding.embed("mock", ["a"], "document");
1166 const u = r as unknown as Uint8Array;
1167 u[0] = 255;
1168 }"#,
1169 Some(provider),
1170 None,
1171 None,
1172 )
1173 .await;
1174 assert!(
1175 cast.result.is_err(),
1176 "casting the result to a Uint8Array must fail, not expose the vectors"
1177 );
1178
1179 for call in [
1184 "crypto.sha256(x)",
1185 "crypto.hmacSha256(crypto.randomBytes(4), x)",
1186 ] {
1187 let provider = MockProvider::new(4, 1_000, Outcome::Reported);
1188 let source = format!(
1189 r#"import embedding from "submilli:embedding";
1190 import crypto from "submilli:crypto";
1191 function main(): void {{
1192 const r = embedding.embed("mock", ["a"], "document");
1193 const x = r as unknown as (string | Uint8Array);
1194 crypto.sha256(crypto.randomBytes(1));
1195 {call};
1196 }}"#
1197 );
1198 let outcome = run(&source, Some(provider), None, None).await;
1199 assert!(
1200 outcome.result.is_err(),
1201 "`{call}` on the result must fail, not hash the vectors"
1202 );
1203 }
1204 }
1205
1206 #[tokio::test]
1210 async fn a_host_read_of_the_result_as_a_uint8array_is_refused() {
1211 let compiled = crate::compile_script(
1212 "function main(): void {}",
1213 "test.ts",
1214 crate::FileId(0),
1215 &[],
1216 &[],
1217 )
1218 .expect("compile clean");
1219 let cfg = RuntimeConfig::default();
1220 let engine = cfg.engine().expect("engine");
1221 let mut data = StoreData::with_vfs(Vfs::tempdir().expect("tempdir"));
1222 data.install_type_info(compiled.type_info.clone());
1223 let mut store = cfg.store_async(&engine, data).expect("store");
1224 crate::runtime::install_tenant_limits(&mut store);
1225 let module = wasmtime::Module::new(&engine, &compiled.wasm).expect("module");
1226 let mut linker = wasmtime::Linker::<StoreData>::new(&engine);
1227 install_runtime_async(&mut linker, &mut store)
1228 .await
1229 .expect("install");
1230 linker
1231 .instantiate_async(&mut store, &module)
1232 .await
1233 .expect("instantiate");
1234
1235 let probe = wasmtime::Func::new(
1236 &mut store,
1237 wasmtime::FuncType::new(&engine, [], []),
1238 |mut caller, _, _| {
1239 let batch = EmbeddingBatch::new(vec![1.0, 2.0], 1, 2, "id", "mock")
1240 .map_err(wasmtime::Error::new)?;
1241 let sealed = super::build_embeddings(&mut caller, batch)?;
1242 let refused =
1243 crate::runtime::host::read_uint8_array_arg(&mut caller, &sealed, "probe");
1244 match refused {
1245 Err(error) if error.to_string().contains("expects a Uint8Array") => Ok(()),
1246 other => Err(wasmtime::Error::msg(format!(
1247 "the sealed result was not refused: {other:?}"
1248 ))),
1249 }
1250 },
1251 );
1252 probe
1253 .call_async(&mut store, &[], &mut [])
1254 .await
1255 .expect("the host refuses to read the sealed result as bytes");
1256 }
1257
1258 #[tokio::test]
1261 async fn the_result_charge_counts_against_the_run_memory_cap() {
1262 let provider = MockProvider::new(3072, 1_000, Outcome::Reported);
1263 let compiled = crate::compile_script(
1264 r#"import embedding from "submilli:embedding";
1265 function main(): void {
1266 const texts: string[] = [];
1267 for (let i = 0; i < 128; i = i + 1) {
1268 texts.push("t");
1269 }
1270 embedding.embed("mock", texts, "document");
1271 }"#,
1272 "test.ts",
1273 crate::FileId(0),
1274 &[],
1275 &[],
1276 )
1277 .expect("compile clean");
1278 let cfg = RuntimeConfig::default();
1279 let engine = cfg.engine().expect("engine");
1280 let mut data = StoreData::with_vfs(Vfs::tempdir().expect("tempdir"));
1281 data.install_type_info(compiled.type_info.clone());
1282 data.embedding_provider = Some(provider);
1283 let budget = budget(EmbeddingLimits::default());
1284 data.embedding_budget = Some(Arc::clone(&budget));
1285 let mut store = cfg.store_async(&engine, data).expect("store");
1286 crate::runtime::install_tenant_limits(&mut store);
1287 let module = wasmtime::Module::new(&engine, &compiled.wasm).expect("module");
1288 let mut linker = wasmtime::Linker::<StoreData>::new(&engine);
1289 install_runtime_async(&mut linker, &mut store)
1290 .await
1291 .expect("install");
1292 let inst = linker
1293 .instantiate_async(&mut store, &module)
1294 .await
1295 .expect("instantiate");
1296 let cap = store.data().tenant_limits.observed_bytes() + 100_000;
1299 store.data_mut().tenant_limits.max_total_bytes = cap;
1300 let result = dispatch_main_async(&mut store, &inst).await;
1301 let error = result.expect_err("the cap refuses the result");
1302 assert!(
1303 crate::runtime::is_memory_exhausted(&error),
1304 "a memory refusal ends the run: {error:?}"
1305 );
1306 assert_eq!(budget.used(), 128, "the reported usage is settled");
1309 assert_eq!(budget.held(), 0, "nothing stays held");
1310 }
1311
1312 #[tokio::test]
1315 async fn an_internal_provider_failure_ends_the_run_and_settles_the_budget() {
1316 let provider = MockProvider::new(4, 1_000, Outcome::Internal);
1317 let budget = budget(EmbeddingLimits::default());
1318 let run = run(
1319 r#"import embedding from "submilli:embedding";
1320 function main(): void {
1321 try {
1322 embedding.embed("mock", ["secret text"], "document");
1323 } catch (e: Error) {
1324 assert(false, "an internal failure must not be catchable");
1325 }
1326 }"#,
1327 Some(provider),
1328 Some(Arc::clone(&budget)),
1329 None,
1330 )
1331 .await;
1332 let error = run.result.expect_err("the run ends");
1333 assert!(
1334 crate::runtime::host::ends_the_run(&error),
1335 "an internal failure is fatal: {error:?}"
1336 );
1337 assert_eq!(budget.held(), 4, "ceil(11 / 3) tokens stay held");
1338 assert_eq!(budget.used(), 4);
1339 }
1340
1341 #[tokio::test]
1344 async fn the_filter_context_carries_the_numbers_and_never_the_text() {
1345 const SECRET: &str = "the patient's diagnosis is confidential";
1346 let provider = MockProvider::new(4, 1_000, Outcome::Reported);
1347 let contexts = Arc::new(Mutex::new(Vec::new()));
1348 let policy: Arc<dyn SecurityCheck> = Arc::new(RecordingPolicy(Arc::clone(&contexts)));
1349 let run = run(
1350 &format!(
1351 r#"import embedding from "submilli:embedding";
1352 function main(): void {{
1353 embedding.embed("mock", ["{SECRET}", "b"], "query");
1354 embedding.models();
1355 }}"#
1356 ),
1357 Some(provider),
1358 None,
1359 Some(policy),
1360 )
1361 .await;
1362 run.result.expect("program completes");
1363
1364 let contexts = contexts.lock().expect("contexts");
1365 assert_eq!(contexts.len(), 2, "one check per embed, one per candidate");
1366 for context in contexts.iter() {
1367 let mut keys: Vec<&str> = context
1368 .as_object()
1369 .expect("an object context")
1370 .keys()
1371 .map(String::as_str)
1372 .collect();
1373 keys.sort_unstable();
1374 assert_eq!(keys, ["input_count", "model"]);
1375 assert!(!context.to_string().contains("patient"), "{context}");
1376 }
1377 assert_eq!(contexts[0]["input_count"], 2, "embed: the text count");
1378 assert_eq!(contexts[1]["input_count"], 0, "models: zero");
1379 }
1380
1381 #[tokio::test]
1384 async fn an_unknown_alias_names_the_available_ones_and_charges_nothing() {
1385 let provider = MockProvider::new(4, 1_000, Outcome::Reported);
1386 let budget = budget(EmbeddingLimits::default());
1387 let run = run(
1388 r#"import embedding from "submilli:embedding";
1389 function main(): void {
1390 let message = "";
1391 try {
1392 embedding.embed("ghost", ["a"], "document");
1393 } catch (e: Error) {
1394 message = e.message;
1395 }
1396 assert(message.indexOf("mock") >= 0, message);
1397 }"#,
1398 Some(Arc::clone(&provider)),
1399 Some(Arc::clone(&budget)),
1400 None,
1401 )
1402 .await;
1403 run.result.expect("program completes");
1404 assert_eq!(provider.calls(), 0);
1405 assert_eq!(budget.used(), 0);
1406 }
1407}