1use std::convert::TryInto;
12use std::fmt;
13use std::sync::{Arc, Mutex};
14use wasmer::wasmparser::{BlockType as WpTypeOrFuncType, Operator};
15use wasmer::{
16 AsStoreMut, ExportIndex, GlobalInit, GlobalType, Instance, LocalFunctionIndex, Mutability,
17 Type,
18 sys::{FunctionMiddleware, MiddlewareError, MiddlewareReaderState, ModuleMiddleware},
19};
20use wasmer_types::{GlobalIndex, ModuleInfo};
21
22#[derive(Clone)]
23struct MeteringGlobalIndexes(GlobalIndex, GlobalIndex);
24
25impl MeteringGlobalIndexes {
26 fn remaining_points(&self) -> GlobalIndex {
28 self.0
29 }
30
31 fn points_exhausted(&self) -> GlobalIndex {
37 self.1
38 }
39}
40
41impl fmt::Debug for MeteringGlobalIndexes {
42 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
43 f.debug_struct("MeteringGlobalIndexes")
44 .field("remaining_points", &self.remaining_points())
45 .field("points_exhausted", &self.points_exhausted())
46 .finish()
47 }
48}
49
50pub struct Metering<F: Fn(&Operator) -> u64 + Send + Sync> {
85 initial_limit: u64,
87
88 cost_function: Arc<F>,
90
91 global_indexes: Mutex<Option<MeteringGlobalIndexes>>,
93}
94
95pub struct FunctionMetering<F: Fn(&Operator) -> u64 + Send + Sync> {
97 cost_function: Arc<F>,
99
100 global_indexes: MeteringGlobalIndexes,
102
103 accumulated_cost: u64,
105}
106
107#[derive(Debug, Eq, PartialEq)]
114pub enum MeteringPoints {
115 Remaining(u64),
119
120 Exhausted,
124}
125
126impl<F: Fn(&Operator) -> u64 + Send + Sync> Metering<F> {
127 pub fn new(initial_limit: u64, cost_function: F) -> Self {
133 Self {
134 initial_limit,
135 cost_function: Arc::new(cost_function),
136 global_indexes: Mutex::new(None),
137 }
138 }
139}
140
141impl<F: Fn(&Operator) -> u64 + Send + Sync> fmt::Debug for Metering<F> {
142 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
143 f.debug_struct("Metering")
144 .field("initial_limit", &self.initial_limit)
145 .field("cost_function", &"<function>")
146 .field("global_indexes", &self.global_indexes)
147 .finish()
148 }
149}
150
151impl<F: Fn(&Operator) -> u64 + Send + Sync + 'static> ModuleMiddleware for Metering<F> {
152 fn generate_function_middleware<'a>(
154 &self,
155 _: LocalFunctionIndex,
156 ) -> Box<dyn FunctionMiddleware<'a> + 'a> {
157 Box::new(FunctionMetering {
158 cost_function: self.cost_function.clone(),
159 global_indexes: self.global_indexes.lock().unwrap().clone().unwrap(),
160 accumulated_cost: 0,
161 })
162 }
163
164 fn transform_module_info(&self, module_info: &mut ModuleInfo) -> Result<(), MiddlewareError> {
166 let mut global_indexes = self.global_indexes.lock().unwrap();
167
168 if global_indexes.is_some() {
169 panic!(
170 "Metering::transform_module_info: Attempting to use a `Metering` middleware from multiple modules."
171 );
172 }
173
174 let remaining_points_global_index = module_info
176 .globals
177 .push(GlobalType::new(Type::I64, Mutability::Var));
178
179 module_info
180 .global_initializers
181 .push(GlobalInit::I64Const(self.initial_limit as i64));
182
183 module_info.exports.insert(
184 "wasmer_metering_remaining_points".to_string(),
185 ExportIndex::Global(remaining_points_global_index),
186 );
187
188 let points_exhausted_global_index = module_info
190 .globals
191 .push(GlobalType::new(Type::I32, Mutability::Var));
192
193 module_info
194 .global_initializers
195 .push(GlobalInit::I32Const(0));
196
197 module_info.exports.insert(
198 "wasmer_metering_points_exhausted".to_string(),
199 ExportIndex::Global(points_exhausted_global_index),
200 );
201
202 *global_indexes = Some(MeteringGlobalIndexes(
203 remaining_points_global_index,
204 points_exhausted_global_index,
205 ));
206
207 Ok(())
208 }
209}
210
211pub fn is_accounting(operator: &Operator) -> bool {
214 matches!(
216 operator,
217 Operator::Loop { .. } | Operator::End | Operator::If { .. } | Operator::Else | Operator::Br { .. } | Operator::BrTable { .. } | Operator::BrIf { .. } | Operator::Call { .. } | Operator::CallIndirect { .. } | Operator::Return | Operator::Throw { .. } | Operator::ThrowRef | Operator::Rethrow { .. } | Operator::Delegate { .. } | Operator::Catch { .. } | Operator::ReturnCall { .. } | Operator::ReturnCallIndirect { .. } | Operator::BrOnCast { .. } | Operator::BrOnCastFail { .. } | Operator::CallRef { .. } | Operator::ReturnCallRef { .. } | Operator::BrOnNull { .. } | Operator::BrOnNonNull { .. } )
245}
246
247impl<F: Fn(&Operator) -> u64 + Send + Sync> fmt::Debug for FunctionMetering<F> {
248 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
249 f.debug_struct("FunctionMetering")
250 .field("cost_function", &"<function>")
251 .field("global_indexes", &self.global_indexes)
252 .finish()
253 }
254}
255
256impl<'a, F: Fn(&Operator) -> u64 + Send + Sync> FunctionMiddleware<'a> for FunctionMetering<F> {
257 fn feed(
258 &mut self,
259 operator: Operator<'a>,
260 state: &mut MiddlewareReaderState<'a>,
261 ) -> Result<(), MiddlewareError> {
262 self.accumulated_cost += (self.cost_function)(&operator);
266
267 if is_accounting(&operator) && self.accumulated_cost > 0 {
269 state.extend(&[
270 Operator::GlobalGet {
272 global_index: self.global_indexes.remaining_points().as_u32(),
273 },
274 Operator::I64Const {
275 value: self.accumulated_cost as i64,
276 },
277 Operator::I64LtU,
278 Operator::If {
279 blockty: WpTypeOrFuncType::Empty,
280 },
281 Operator::I32Const { value: 1 },
282 Operator::GlobalSet {
283 global_index: self.global_indexes.points_exhausted().as_u32(),
284 },
285 Operator::Unreachable,
286 Operator::End,
287 Operator::GlobalGet {
289 global_index: self.global_indexes.remaining_points().as_u32(),
290 },
291 Operator::I64Const {
292 value: self.accumulated_cost as i64,
293 },
294 Operator::I64Sub,
295 Operator::GlobalSet {
296 global_index: self.global_indexes.remaining_points().as_u32(),
297 },
298 ]);
299
300 self.accumulated_cost = 0;
301 }
302 state.push_operator(operator);
303
304 Ok(())
305 }
306}
307
308pub fn get_remaining_points(ctx: &mut impl AsStoreMut, instance: &Instance) -> MeteringPoints {
333 let exhausted: i32 = instance
334 .exports
335 .get_global("wasmer_metering_points_exhausted")
336 .expect("Can't get `wasmer_metering_points_exhausted` from Instance")
337 .get(ctx)
338 .try_into()
339 .expect("`wasmer_metering_points_exhausted` from Instance has wrong type");
340
341 if exhausted > 0 {
342 return MeteringPoints::Exhausted;
343 }
344
345 let points = instance
346 .exports
347 .get_global("wasmer_metering_remaining_points")
348 .expect("Can't get `wasmer_metering_remaining_points` from Instance")
349 .get(ctx)
350 .try_into()
351 .expect("`wasmer_metering_remaining_points` from Instance has wrong type");
352
353 MeteringPoints::Remaining(points)
354}
355
356pub fn set_remaining_points(ctx: &mut impl AsStoreMut, instance: &Instance, points: u64) {
382 instance
383 .exports
384 .get_global("wasmer_metering_remaining_points")
385 .expect("Can't get `wasmer_metering_remaining_points` from Instance")
386 .set(ctx, points.into())
387 .expect("Can't set `wasmer_metering_remaining_points` in Instance");
388
389 instance
390 .exports
391 .get_global("wasmer_metering_points_exhausted")
392 .expect("Can't get `wasmer_metering_points_exhausted` from Instance")
393 .set(ctx, 0i32.into())
394 .expect("Can't set `wasmer_metering_points_exhausted` in Instance");
395}
396
397#[cfg(all(test, not(target_os = "windows")))]
400mod tests {
401 use super::*;
402
403 use std::sync::Arc;
404 use wasmer::sys::EngineBuilder;
405 use wasmer::{
406 Module, Store, TypedFunction, imports,
407 sys::{CompilerConfig, Cranelift},
408 wat2wasm,
409 };
410
411 fn cost_function(operator: &Operator) -> u64 {
412 match operator {
413 Operator::LocalGet { .. } | Operator::I32Const { .. } => 1,
414 Operator::I32Add { .. } => 2,
415 _ => 0,
416 }
417 }
418
419 fn bytecode() -> Vec<u8> {
420 wat2wasm(
421 br#"(module
422 (type $add_t (func (param i32) (result i32)))
423 (func $add_one_f (type $add_t) (param $value i32) (result i32)
424 local.get $value
425 i32.const 1
426 i32.add)
427 (func $short_loop_f
428 (local $x f64) (local $j i32)
429 (local.set $x (f64.const 5.5))
430
431 (loop $named_loop
432 ;; $j++
433 local.get $j
434 i32.const 1
435 i32.add
436 local.set $j
437
438 ;; if $j < 5, one more time
439 local.get $j
440 i32.const 5
441 i32.lt_s
442 br_if $named_loop
443 )
444 )
445 (func $infi_loop_f
446 (loop $infi_loop_start
447 br $infi_loop_start
448 )
449 )
450 (export "add_one" (func $add_one_f))
451 (export "short_loop" (func $short_loop_f))
452 (export "infi_loop" (func $infi_loop_f))
453 )"#,
454 )
455 .unwrap()
456 .into()
457 }
458
459 #[test]
460 fn get_remaining_points_works() {
461 let metering = Arc::new(Metering::new(10, cost_function));
462 let mut compiler_config = Cranelift::default();
463 compiler_config.push_middleware(metering);
464 let mut store = Store::new(EngineBuilder::new(compiler_config));
465 let module = Module::new(&store, bytecode()).unwrap();
466
467 let instance = Instance::new(&mut store, &module, &imports! {}).unwrap();
469 assert_eq!(
470 get_remaining_points(&mut store, &instance),
471 MeteringPoints::Remaining(10)
472 );
473
474 let add_one: TypedFunction<i32, i32> = instance
481 .exports
482 .get_function("add_one")
483 .unwrap()
484 .typed(&store)
485 .unwrap();
486 add_one.call(&mut store, 1).unwrap();
487 assert_eq!(
488 get_remaining_points(&mut store, &instance),
489 MeteringPoints::Remaining(6)
490 );
491
492 add_one.call(&mut store, 1).unwrap();
494 assert_eq!(
495 get_remaining_points(&mut store, &instance),
496 MeteringPoints::Remaining(2)
497 );
498
499 assert!(add_one.call(&mut store, 1).is_err());
501 assert_eq!(
502 get_remaining_points(&mut store, &instance),
503 MeteringPoints::Exhausted
504 );
505 }
506
507 #[test]
508 fn set_remaining_points_works() {
509 let metering = Arc::new(Metering::new(10, cost_function));
510 let mut compiler_config = Cranelift::default();
511 compiler_config.push_middleware(metering);
512 let mut store = Store::new(EngineBuilder::new(compiler_config));
513 let module = Module::new(&store, bytecode()).unwrap();
514
515 let instance = Instance::new(&mut store, &module, &imports! {}).unwrap();
517 assert_eq!(
518 get_remaining_points(&mut store, &instance),
519 MeteringPoints::Remaining(10)
520 );
521 let add_one: TypedFunction<i32, i32> = instance
522 .exports
523 .get_function("add_one")
524 .unwrap()
525 .typed(&store)
526 .unwrap();
527
528 set_remaining_points(&mut store, &instance, 12);
530
531 add_one.call(&mut store, 1).unwrap();
533 assert_eq!(
534 get_remaining_points(&mut store, &instance),
535 MeteringPoints::Remaining(8)
536 );
537
538 add_one.call(&mut store, 1).unwrap();
539 assert_eq!(
540 get_remaining_points(&mut store, &instance),
541 MeteringPoints::Remaining(4)
542 );
543
544 add_one.call(&mut store, 1).unwrap();
545 assert_eq!(
546 get_remaining_points(&mut store, &instance),
547 MeteringPoints::Remaining(0)
548 );
549
550 assert!(add_one.call(&mut store, 1).is_err());
551 assert_eq!(
552 get_remaining_points(&mut store, &instance),
553 MeteringPoints::Exhausted
554 );
555
556 set_remaining_points(&mut store, &instance, 4);
558 assert_eq!(
559 get_remaining_points(&mut store, &instance),
560 MeteringPoints::Remaining(4)
561 );
562 }
563
564 #[test]
565 fn metering_works_for_loops() {
566 const INITIAL_POINTS: u64 = 10_000;
567
568 fn cost(operator: &Operator) -> u64 {
569 match operator {
570 Operator::Loop { .. } => 1000,
571 Operator::Br { .. } | Operator::BrIf { .. } => 10,
572 Operator::F64Const { .. } => 7,
573 _ => 0,
574 }
575 }
576
577 let metering = Arc::new(Metering::new(INITIAL_POINTS, cost));
580 let mut compiler_config = Cranelift::default();
581 compiler_config.push_middleware(metering);
582 let mut store = Store::new(EngineBuilder::new(compiler_config));
583 let module = Module::new(&store, bytecode()).unwrap();
584
585 let instance = Instance::new(&mut store, &module, &imports! {}).unwrap();
586 let short_loop: TypedFunction<(), ()> = instance
587 .exports
588 .get_function("short_loop")
589 .unwrap()
590 .typed(&store)
591 .unwrap();
592 short_loop.call(&mut store).unwrap();
593
594 let points_used: u64 = match get_remaining_points(&mut store, &instance) {
595 MeteringPoints::Exhausted => panic!("Unexpected exhausted"),
596 MeteringPoints::Remaining(remaining) => INITIAL_POINTS - remaining,
597 };
598
599 assert_eq!(
600 points_used,
601 7 +
602 1000 + 50 );
604
605 let metering = Arc::new(Metering::new(INITIAL_POINTS, cost));
608 let mut compiler_config = Cranelift::default();
609 compiler_config.push_middleware(metering);
610 let mut store = Store::new(EngineBuilder::new(compiler_config));
611 let module = Module::new(&store, bytecode()).unwrap();
612
613 let instance = Instance::new(&mut store, &module, &imports! {}).unwrap();
614 let infi_loop: TypedFunction<(), ()> = instance
615 .exports
616 .get_function("infi_loop")
617 .unwrap()
618 .typed(&store)
619 .unwrap();
620 infi_loop.call(&mut store).unwrap_err(); assert_eq!(
623 get_remaining_points(&mut store, &instance),
624 MeteringPoints::Exhausted
625 );
626 }
627}