LuaDisplay.display_fn
change rwrvylpxvtsutturlpplquwzuznozklx
commit f48e5ad0b5faea8ef9fa230de15498a47e4d537b
author Alpha Chen <alpha@kejadlen.dev>
date
parent wrlntwsr
diff --git a/src/main.rs b/src/main.rs
index 7e3175c..31336ca 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -187,14 +187,21 @@ struct LuaAssertion {
     name: String,
     args: LuaMultiValue,
     status_fn: LuaFunction,
+    display_fn: Option<LuaFunction>,
 }
 
 impl LuaAssertion {
-    fn new(name: &str, args: LuaMultiValue, status_fn: LuaFunction) -> Self {
+    fn new(
+        name: &str,
+        args: LuaMultiValue,
+        status_fn: LuaFunction,
+        display_fn: Option<LuaFunction>,
+    ) -> Self {
         Self {
             name: name.to_string(),
             args,
             status_fn,
+            display_fn,
         }
     }
 }
@@ -213,6 +220,18 @@ impl AssertionType for LuaAssertion {
     }
 
     fn display(&self) -> String {
+        if let Some(ref display_fn) = self.display_fn {
+            display_fn
+                .call::<String>(self.args.clone())
+                .unwrap_or_else(|_| self.default_display())
+        } else {
+            self.default_display()
+        }
+    }
+}
+
+impl LuaAssertion {
+    fn default_display(&self) -> String {
         let args_str = self
             .args
             .iter()
@@ -240,6 +259,7 @@ impl Default for Frork {
 impl Frork {
     fn register(&self, name: &str, table: LuaTable) -> LuaResult<()> {
         let status_fn: LuaFunction = table.get("status")?;
+        let display_fn: Option<LuaFunction> = table.get("display").ok();
 
         let name_clone = name.to_string();
         self.registry.borrow_mut().register(name, move |args| {
@@ -247,6 +267,7 @@ impl Frork {
                 &name_clone,
                 args,
                 status_fn.clone(),
+                display_fn.clone(),
             )))
         });
         info!("Registered assertion type: {}", name);
@@ -352,7 +373,7 @@ mod tests {
             .create_function(|_lua, _args: LuaMultiValue| Ok("ok".to_string()))
             .unwrap();
 
-        let assertion = LuaAssertion::new("test", LuaMultiValue::new(), status_fn);
+        let assertion = LuaAssertion::new("test", LuaMultiValue::new(), status_fn, None);
         let result = assertion.status().unwrap();
 
         match result {
@@ -368,7 +389,7 @@ mod tests {
             .create_function(|_lua, _args: LuaMultiValue| Ok("missing".to_string()))
             .unwrap();
 
-        let assertion = LuaAssertion::new("test", LuaMultiValue::new(), status_fn);
+        let assertion = LuaAssertion::new("test", LuaMultiValue::new(), status_fn, None);
         let result = assertion.status().unwrap();
 
         match result {
@@ -384,7 +405,7 @@ mod tests {
             .create_function(|_lua, _args: LuaMultiValue| Ok("invalid".to_string()))
             .unwrap();
 
-        let assertion = LuaAssertion::new("test", LuaMultiValue::new(), status_fn);
+        let assertion = LuaAssertion::new("test", LuaMultiValue::new(), status_fn, None);
         let result = assertion.status();
 
         assert!(result.is_err());
@@ -400,10 +421,36 @@ mod tests {
         let args = vec![LuaValue::String(lua.create_string("arg1").unwrap())];
         let lua_args = LuaMultiValue::from_vec(args);
 
-        let assertion = LuaAssertion::new("mytest", lua_args, status_fn);
+        let assertion = LuaAssertion::new("mytest", lua_args, status_fn, None);
         let display = assertion.display();
 
         assert!(display.contains("mytest"));
         assert!(display.contains("arg1"));
     }
+
+    #[test]
+    fn test_lua_assertion_custom_display() {
+        let lua = Lua::new();
+        let status_fn = lua
+            .create_function(|_lua, _args: LuaMultiValue| Ok("ok".to_string()))
+            .unwrap();
+        let display_fn = lua
+            .create_function(|_lua, args: LuaMultiValue| {
+                let arg_str = args
+                    .into_iter()
+                    .map(|v| v.to_string().unwrap_or_else(|_| "?".to_string()))
+                    .collect::<Vec<_>>()
+                    .join(",");
+                Ok(format!("custom: {}", arg_str))
+            })
+            .unwrap();
+
+        let args = vec![LuaValue::String(lua.create_string("test1").unwrap())];
+        let lua_args = LuaMultiValue::from_vec(args);
+
+        let assertion = LuaAssertion::new("mytest", lua_args, status_fn, Some(display_fn));
+        let display = assertion.display();
+
+        assert_eq!(display, "custom: test1");
+    }
 }