diff --git a/.gitignore b/.gitignore index b5885de..9634914 100644 --- a/.gitignore +++ b/.gitignore @@ -25,6 +25,12 @@ bindings/ffi/regorus.ffi.hpp bindings/*/target +# Temporary commit message files +.commit-msg.txt + +# Local planning docs +docs/plans/ + # C# build folders **bin **obj diff --git a/Cargo.lock b/Cargo.lock index a0323d5..15f3414 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -817,6 +817,12 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "lru" +version = "0.16.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1dc47f592c06f33f8e3aea9591776ec7c9f9e4124778ff8a3c3b87159f7e593" + [[package]] name = "memchr" version = "2.7.6" @@ -1265,10 +1271,12 @@ dependencies = [ "ipnet", "jsonschema", "lazy_static", + "lru", "msvc_spectre_libs", "num-bigint", "num-traits", "num_cpus", + "parking_lot", "postcard", "prettydiff", "rand", diff --git a/Cargo.toml b/Cargo.toml index fb45997..2dbe5d3 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -39,10 +39,11 @@ net = ["dep:ipnet"] no_std = ["lazy_static/spin_no_std"] opa-runtime = [] regex = ["dep:regex"] +cache = ["dep:lru"] rvm = ["dep:postcard", "dep:indexmap"] semver = ["dep:semver"] allocator-memory-limits = ["std", "mimalloc", "mimalloc/allocator-memory-limits"] -std = ["rand/std", "rand/std_rng", "serde_json/std", "msvc_spectre_libs" ] +std = ["rand/std", "rand/std_rng", "serde_json/std", "msvc_spectre_libs", "dep:parking_lot" ] time = ["dep:chrono", "dep:chrono-tz"] uuid = ["dep:uuid"] urlquery = ["dep:url"] @@ -61,6 +62,7 @@ full-opa = [ "net", "opa-runtime", "regex", + "cache", "semver", "std", "time", @@ -105,6 +107,7 @@ thiserror = { version = "2.0", default-features = false } data-encoding = { version = "2.8.0", optional = true, default-features=false, features = ["alloc"] } num-bigint = { version = "0.4", default-features = false } num-traits = { version = "0.2", default-features = false } +parking_lot = { version = "0.12", optional = true } spin = { version = "0.9.8", default-features = false, features = ["mutex", "spin_mutex"] } globset = { version = "0.4.16", features = ["simd-accel"], default-features = false, optional = true } @@ -124,6 +127,7 @@ rand = { version = "0.9.0", default-features = false, features = ["thread_rng"], # Causes the project to link with the Spectre-mitigated CRT and libs. msvc_spectre_libs = { version = "0.1", features = ["error"], optional = true } dashmap = { version = "6.1", default-features = false, optional = true } +lru = { version = "0.16", default-features = false, optional = true } mimalloc = { package = "regorus-mimalloc", path = "mimalloc", version = "2.2.6", optional = true } # rvm related deps diff --git a/bindings/c/main.c b/bindings/c/main.c index e9f01ee..944325f 100644 --- a/bindings/c/main.c +++ b/bindings/c/main.c @@ -11,6 +11,13 @@ int main() { if (r.status != Ok) goto error; + // Configure the global pattern caches. + RegorusCacheConfig cache_config = { .regex = 256, .glob = 128 }; + r = regorus_set_cache_config(cache_config); + if (r.status != Ok) + goto error; + regorus_result_drop(r); + // Raise the default col limit to 2000 RegorusPolicyLengthConfig len_config = { .max_col = 2000, .max_file_bytes = 1048576, .max_lines = 20000 }; r = regorus_engine_set_policy_length_config(engine, len_config); diff --git a/bindings/cpp/main.cpp b/bindings/cpp/main.cpp index 63cd92e..0ea7923 100644 --- a/bindings/cpp/main.cpp +++ b/bindings/cpp/main.cpp @@ -6,6 +6,10 @@ void example() // Create engine regorus::Engine engine; + // Configure the global pattern caches. + RegorusCacheConfig cache_config = { 256, 128 }; + regorus::set_cache_config(cache_config); + engine.set_rego_v0(true); engine.set_enable_coverage(true); diff --git a/bindings/cpp/regorus.hpp b/bindings/cpp/regorus.hpp index 5d43d70..191e4cf 100644 --- a/bindings/cpp/regorus.hpp +++ b/bindings/cpp/regorus.hpp @@ -158,6 +158,14 @@ namespace regorus { Engine& operator=(const Engine&) = delete; }; + inline Result set_cache_config(RegorusCacheConfig config) { + return Result(regorus_set_cache_config(config)); + } + + inline Result clear_cache() { + return Result(regorus_clear_cache()); + } + class CompiledPolicy { public: explicit CompiledPolicy(RegorusCompiledPolicy* p) : policy(p) {} diff --git a/bindings/csharp/Regorus/CacheConfig.cs b/bindings/csharp/Regorus/CacheConfig.cs new file mode 100644 index 0000000..692a61e --- /dev/null +++ b/bindings/csharp/Regorus/CacheConfig.cs @@ -0,0 +1,39 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System; + +namespace Regorus +{ + /// + /// Global configuration for compiled pattern caches used by regex and glob builtins. + /// + public readonly struct CacheConfig + { + /// + /// Initializes a new instance of the struct. + /// + /// Maximum cached compiled regex patterns (default 256, 0 = disabled). + /// Maximum cached compiled glob matchers (default 128, 0 = disabled). + public CacheConfig(nuint regex, nuint glob) + { + Regex = regex; + Glob = glob; + } + + /// Maximum cached compiled regex patterns (default 256). + public nuint Regex { get; } + + /// Maximum cached compiled glob matchers (default 128). + public nuint Glob { get; } + + internal Regorus.Internal.RegorusCacheConfig ToNative() + { + return new Regorus.Internal.RegorusCacheConfig + { + regex = Regex, + glob = Glob, + }; + } + } +} diff --git a/bindings/csharp/Regorus/Engine.cs b/bindings/csharp/Regorus/Engine.cs index 78d14de..852be97 100644 --- a/bindings/csharp/Regorus/Engine.cs +++ b/bindings/csharp/Regorus/Engine.cs @@ -34,6 +34,17 @@ namespace Regorus CheckAndDropResult(Regorus.Internal.API.regorus_clear_fallback_execution_timer_config()); } + public static void SetCacheConfig(CacheConfig config) + { + var nativeConfig = config.ToNative(); + CheckAndDropResult(Regorus.Internal.API.regorus_set_cache_config(nativeConfig)); + } + + public static void ClearCache() + { + CheckAndDropResult(Regorus.Internal.API.regorus_clear_cache()); + } + private Engine(RegorusEngineHandle handle) : base(handle, nameof(Engine)) { diff --git a/bindings/csharp/Regorus/NativeMethods.cs b/bindings/csharp/Regorus/NativeMethods.cs index df2170d..6702d7f 100644 --- a/bindings/csharp/Regorus/NativeMethods.cs +++ b/bindings/csharp/Regorus/NativeMethods.cs @@ -458,6 +458,22 @@ namespace Regorus.Internal #endregion + #region Cache Configuration Global Methods + + /// + /// Configure the global pattern caches used by regex and glob builtins. + /// + [DllImport(LibraryName, EntryPoint = "regorus_set_cache_config", CallingConvention = CallingConvention.Cdecl, ExactSpelling = true)] + internal static extern RegorusResult regorus_set_cache_config(RegorusCacheConfig config); + + /// + /// Clear all entries from every pattern cache. + /// + [DllImport(LibraryName, EntryPoint = "regorus_clear_cache", CallingConvention = CallingConvention.Cdecl, ExactSpelling = true)] + internal static extern RegorusResult regorus_clear_cache(); + + #endregion + #region Compilation Methods /// @@ -795,6 +811,16 @@ namespace Regorus.Internal public UIntPtr max_lines; } + /// + /// FFI representation of the cache configuration. + /// + [StructLayout(LayoutKind.Sequential)] + internal struct RegorusCacheConfig + { + public UIntPtr regex; + public UIntPtr glob; + } + /// /// Byte buffer returned from FFI. /// diff --git a/bindings/csharp/TestApp/Program.cs b/bindings/csharp/TestApp/Program.cs index 1cfb42f..ebd0b47 100644 --- a/bindings/csharp/TestApp/Program.cs +++ b/bindings/csharp/TestApp/Program.cs @@ -18,6 +18,9 @@ var w = new Stopwatch(); w.Restart(); +// Configure the global pattern caches. +Regorus.Engine.SetCacheConfig(new Regorus.CacheConfig(regex: 256, glob: 128)); + var engine = new Regorus.Engine(); engine.SetRegoV0(true); // Raise the default col limit to 2000 diff --git a/bindings/ffi/Cargo.lock b/bindings/ffi/Cargo.lock index dec168c..698a6bc 100644 --- a/bindings/ffi/Cargo.lock +++ b/bindings/ffi/Cargo.lock @@ -657,6 +657,12 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "lru" +version = "0.16.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1dc47f592c06f33f8e3aea9591776ec7c9f9e4124778ff8a3c3b87159f7e593" + [[package]] name = "memchr" version = "2.7.6" @@ -985,9 +991,11 @@ dependencies = [ "ipnet", "jsonschema", "lazy_static", + "lru", "msvc_spectre_libs", "num-bigint", "num-traits", + "parking_lot", "postcard", "rand", "regex", diff --git a/bindings/ffi/Cargo.toml b/bindings/ffi/Cargo.toml index a787910..f20c3f7 100644 --- a/bindings/ffi/Cargo.toml +++ b/bindings/ffi/Cargo.toml @@ -35,6 +35,7 @@ default = [ "rbac", "regorus/arc", "regorus/full-opa", + "cache", "contention_checks", ] ast = ["regorus/ast"] @@ -45,6 +46,7 @@ allocator-memory-limits = ["regorus/allocator-memory-limits"] contention_checks = ["parking_lot"] rvm = ["regorus/rvm"] rbac = ["regorus/azure-rbac"] +cache = ["regorus/cache"] custom_allocator = [] [build-dependencies] diff --git a/bindings/ffi/src/limits.rs b/bindings/ffi/src/limits.rs index 18b74ff..dc98a4e 100644 --- a/bindings/ffi/src/limits.rs +++ b/bindings/ffi/src/limits.rs @@ -199,6 +199,40 @@ pub extern "C" fn regorus_clear_fallback_execution_timer_config() -> RegorusResu RegorusResult::ok_void() } +// --------------------------------------------------------------------------- +// Cache configuration (global) +// --------------------------------------------------------------------------- + +/// FFI representation of [`regorus::cache::Config`]. +#[cfg(feature = "cache")] +#[repr(C)] +#[derive(Debug, Clone, Copy)] +pub struct RegorusCacheConfig { + /// Maximum compiled regex patterns (default 256, 0 = disabled). + pub regex: usize, + /// Maximum compiled glob matchers (default 128, 0 = disabled). + pub glob: usize, +} + +/// Configure the global pattern caches used by `regex.*` and `glob.*` builtins. +#[cfg(feature = "cache")] +#[no_mangle] +pub extern "C" fn regorus_set_cache_config(config: RegorusCacheConfig) -> RegorusResult { + regorus::cache::configure(regorus::cache::Config { + regex: config.regex, + glob: config.glob, + }); + RegorusResult::ok_void() +} + +/// Clear all entries from every pattern cache. +#[cfg(feature = "cache")] +#[no_mangle] +pub extern "C" fn regorus_clear_cache() -> RegorusResult { + regorus::cache::clear(); + RegorusResult::ok_void() +} + #[cfg(test)] mod tests { use super::{ diff --git a/bindings/go/main.go b/bindings/go/main.go index e59de0d..a5d4ed3 100644 --- a/bindings/go/main.go +++ b/bindings/go/main.go @@ -17,6 +17,12 @@ func main() { engine := regorus.NewEngine() defer engine.Close() + // Configure the global pattern caches. + if err = regorus.SetCacheConfig(regorus.CacheConfig{Regex: 256, Glob: 128}); err != nil { + fmt.Fprintf(os.Stderr, "error: %v\n", err) + os.Exit(1) + } + engine.SetRegoV0(true) // Raise the default col limit to 2000 engine.SetPolicyLengthConfig(regorus.PolicyLengthConfig{MaxCol: 2000, MaxFileBytes: 1048576, MaxLines: 20000}) diff --git a/bindings/go/pkg/regorus/mod.go b/bindings/go/pkg/regorus/mod.go index 98aaf49..042f2b3 100644 --- a/bindings/go/pkg/regorus/mod.go +++ b/bindings/go/pkg/regorus/mod.go @@ -243,3 +243,30 @@ func (e *Engine) ClearPolicyLengthConfig() error { } return nil } + +type CacheConfig struct { + Regex uint + Glob uint +} + +func SetCacheConfig(config CacheConfig) error { + c := C.RegorusCacheConfig{ + regex: C.size_t(config.Regex), + glob: C.size_t(config.Glob), + } + result := C.regorus_set_cache_config(c) + defer C.regorus_result_drop(result) + if result.status != C.Ok { + return fmt.Errorf("%s", C.GoString(result.error_message)) + } + return nil +} + +func ClearCache() error { + result := C.regorus_clear_cache() + defer C.regorus_result_drop(result) + if result.status != C.Ok { + return fmt.Errorf("%s", C.GoString(result.error_message)) + } + return nil +} diff --git a/bindings/java/Cargo.lock b/bindings/java/Cargo.lock index 280df93..91ae70e 100644 --- a/bindings/java/Cargo.lock +++ b/bindings/java/Cargo.lock @@ -539,6 +539,12 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "lru" +version = "0.16.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1dc47f592c06f33f8e3aea9591776ec7c9f9e4124778ff8a3c3b87159f7e593" + [[package]] name = "memchr" version = "2.7.6" @@ -860,9 +866,11 @@ dependencies = [ "ipnet", "jsonschema", "lazy_static", + "lru", "msvc_spectre_libs", "num-bigint", "num-traits", + "parking_lot", "postcard", "rand", "regex", diff --git a/bindings/java/Cargo.toml b/bindings/java/Cargo.toml index dc6d44a..0710448 100644 --- a/bindings/java/Cargo.toml +++ b/bindings/java/Cargo.toml @@ -14,9 +14,10 @@ keywords = ["interpreter", "opa", "policy-as-code", "rego"] crate-type = ["cdylib"] [features] -default = ["ast", "coverage", "regorus/std", "regorus/full-opa"] +default = ["ast", "cache", "coverage", "regorus/std", "regorus/full-opa"] coverage = ["regorus/coverage"] ast = ["regorus/ast"] +cache = ["regorus/cache"] [dependencies] anyhow = "1.0" diff --git a/bindings/java/Test.java b/bindings/java/Test.java index 08f2371..5f802a9 100644 --- a/bindings/java/Test.java +++ b/bindings/java/Test.java @@ -1,6 +1,7 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. +import com.microsoft.regorus.CacheConfig; import com.microsoft.regorus.Engine; import com.microsoft.regorus.PolicyLengthConfig; import com.microsoft.regorus.PolicyModule; @@ -10,6 +11,9 @@ import com.microsoft.regorus.Rvm; public class Test { public static void main(String[] args) { + // Configure the global pattern caches. + CacheConfig.configure(new CacheConfig(256, 128)); + try (Engine engine = new Engine()) { String pkg = engine.addPolicy( "hello.rego", diff --git a/bindings/java/src/lib.rs b/bindings/java/src/lib.rs index 6732c39..7f6a1a5 100644 --- a/bindings/java/src/lib.rs +++ b/bindings/java/src/lib.rs @@ -399,6 +399,37 @@ pub extern "system" fn Java_com_microsoft_regorus_Engine_nativeClearPolicyLength engine.clear_policy_length_config(); } +#[cfg(feature = "cache")] +#[no_mangle] +pub extern "system" fn Java_com_microsoft_regorus_CacheConfig_nativeSetCacheConfig( + _env: JNIEnv, + _class: JClass, + regex: jlong, + glob: jlong, +) { + regorus::cache::configure(regorus::cache::Config { + regex: if regex < 0 { + 0 + } else { + usize::try_from(regex).unwrap_or(usize::MAX) + }, + glob: if glob < 0 { + 0 + } else { + usize::try_from(glob).unwrap_or(usize::MAX) + }, + }); +} + +#[cfg(feature = "cache")] +#[no_mangle] +pub extern "system" fn Java_com_microsoft_regorus_CacheConfig_nativeClearCache( + _env: JNIEnv, + _class: JClass, +) { + regorus::cache::clear(); +} + #[no_mangle] pub extern "system" fn Java_com_microsoft_regorus_Engine_nativeDestroyEngine( _env: JNIEnv, diff --git a/bindings/java/src/main/java/com/microsoft/regorus/CacheConfig.java b/bindings/java/src/main/java/com/microsoft/regorus/CacheConfig.java new file mode 100644 index 0000000..06354d2 --- /dev/null +++ b/bindings/java/src/main/java/com/microsoft/regorus/CacheConfig.java @@ -0,0 +1,62 @@ +/** + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT License. + **/ + +package com.microsoft.regorus; + +/** + * Global configuration for compiled pattern caches used by regex and glob builtins. + * + *

Capacity of 0 disables the corresponding cache. + */ +public final class CacheConfig { + + static { + System.loadLibrary("regorus_java"); + } + + private static native void nativeSetCacheConfig(long regex, long glob); + private static native void nativeClearCache(); + + /** + * Maximum cached compiled regex patterns (default 256). + */ + public final long regex; + + /** + * Maximum cached compiled glob matchers (default 128). + */ + public final long glob; + + /** + * Create a new cache configuration. + * + * @param regex Maximum cached compiled regex patterns (0 = disabled). + * @param glob Maximum cached compiled glob matchers (0 = disabled). + */ + public CacheConfig(long regex, long glob) { + if (regex < 0) { + throw new IllegalArgumentException("regex must be non-negative"); + } + if (glob < 0) { + throw new IllegalArgumentException("glob must be non-negative"); + } + this.regex = regex; + this.glob = glob; + } + + /** + * Apply this cache configuration globally. + */ + public static void configure(CacheConfig config) { + nativeSetCacheConfig(config.regex, config.glob); + } + + /** + * Clear all entries from every pattern cache. + */ + public static void clear() { + nativeClearCache(); + } +} diff --git a/bindings/python/Cargo.lock b/bindings/python/Cargo.lock index d4c819f..4e5798b 100644 --- a/bindings/python/Cargo.lock +++ b/bindings/python/Cargo.lock @@ -510,6 +510,12 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "lru" +version = "0.16.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1dc47f592c06f33f8e3aea9591776ec7c9f9e4124778ff8a3c3b87159f7e593" + [[package]] name = "memchr" version = "2.7.6" @@ -919,9 +925,11 @@ dependencies = [ "ipnet", "jsonschema", "lazy_static", + "lru", "msvc_spectre_libs", "num-bigint", "num-traits", + "parking_lot", "postcard", "rand", "regex", diff --git a/bindings/python/Cargo.toml b/bindings/python/Cargo.toml index 1d721e9..87f2c94 100644 --- a/bindings/python/Cargo.toml +++ b/bindings/python/Cargo.toml @@ -15,8 +15,9 @@ keywords = ["interpreter", "opa", "policy-as-code", "rego"] crate-type = ["cdylib"] [features] -default = ["ast", "coverage", "regorus/std", "regorus/full-opa"] +default = ["ast", "cache", "coverage", "regorus/std", "regorus/full-opa"] ast = ["regorus/ast"] +cache = ["regorus/cache"] coverage = ["regorus/coverage"] [dependencies] diff --git a/bindings/python/src/lib.rs b/bindings/python/src/lib.rs index e97362b..bff5061 100644 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -622,10 +622,33 @@ impl Rvm { } } +/// Configure the global pattern caches used by `regex.*` and `glob.*` builtins. +/// +/// * `regex`: Maximum cached compiled regex patterns (default 256, 0 = disabled). +/// * `glob`: Maximum cached compiled glob matchers (default 128, 0 = disabled). +#[cfg(feature = "cache")] +#[pyfunction] +#[pyo3(signature = (*, regex = 256, glob = 128))] +fn set_cache_config(regex: usize, glob: usize) { + ::regorus::cache::configure(::regorus::cache::Config { regex, glob }); +} + +/// Clear all entries from every pattern cache. +#[cfg(feature = "cache")] +#[pyfunction] +fn clear_cache() { + ::regorus::cache::clear(); +} + #[pymodule] pub fn regorus(_py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; + #[cfg(feature = "cache")] + { + m.add_function(wrap_pyfunction!(set_cache_config, m)?)?; + m.add_function(wrap_pyfunction!(clear_cache, m)?)?; + } Ok(()) } diff --git a/bindings/python/test.py b/bindings/python/test.py index 2f6262e..5ef1bef 100644 --- a/bindings/python/test.py +++ b/bindings/python/test.py @@ -6,6 +6,9 @@ import sys if hasattr(sys.stdout, "reconfigure"): sys.stdout.reconfigure(encoding="utf-8") +# Configure the global pattern caches. +regorus.set_cache_config(regex=256, glob=128) + # Create engine engine = regorus.Engine() diff --git a/bindings/ruby/Cargo.lock b/bindings/ruby/Cargo.lock index 92d4882..b46e74b 100644 --- a/bindings/ruby/Cargo.lock +++ b/bindings/ruby/Cargo.lock @@ -549,6 +549,12 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "lru" +version = "0.16.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1dc47f592c06f33f8e3aea9591776ec7c9f9e4124778ff8a3c3b87159f7e593" + [[package]] name = "magnus" version = "0.8.2" @@ -926,6 +932,7 @@ dependencies = [ "ipnet", "jsonschema", "lazy_static", + "lru", "msvc_spectre_libs", "num-bigint", "num-traits", diff --git a/bindings/ruby/ext/regorusrb/Cargo.toml b/bindings/ruby/ext/regorusrb/Cargo.toml index 7d7de82..b2f1851 100644 --- a/bindings/ruby/ext/regorusrb/Cargo.toml +++ b/bindings/ruby/ext/regorusrb/Cargo.toml @@ -11,8 +11,9 @@ crate-type = ["cdylib"] path = "src/lib.rs" [features] -default = ["ast", "coverage", "regorus/std", "regorus/full-opa"] +default = ["ast", "cache", "coverage", "regorus/std", "regorus/full-opa"] ast = ["regorus/ast"] +cache = ["regorus/cache"] coverage = ["regorus/coverage"] [dependencies] diff --git a/bindings/ruby/ext/regorusrb/src/lib.rs b/bindings/ruby/ext/regorusrb/src/lib.rs index c4142a4..c1746e4 100644 --- a/bindings/ruby/ext/regorusrb/src/lib.rs +++ b/bindings/ruby/ext/regorusrb/src/lib.rs @@ -14,6 +14,13 @@ struct PolicyLengthSpec { max_lines: usize, } +#[cfg(feature = "cache")] +#[derive(Deserialize)] +struct CacheConfigSpec { + regex: usize, + glob: usize, +} + #[derive(Default)] #[magnus::wrap(class = "Regorus::Engine")] pub struct Engine { @@ -417,5 +424,35 @@ fn init(ruby: &Ruby) -> Result<(), Error> { // ast engine_class.define_method("get_ast_as_json", method!(Engine::get_ast_as_json, 0))?; + + // cache configuration (module-level) + #[cfg(feature = "cache")] + { + regorus_module + .define_module_function("set_cache_config", magnus::function!(set_cache_config, 1))?; + regorus_module.define_module_function("clear_cache", magnus::function!(clear_cache, 0))?; + } + + Ok(()) +} + +#[cfg(feature = "cache")] +fn set_cache_config(ruby: &Ruby, hash: magnus::RHash) -> Result<(), Error> { + let spec: CacheConfigSpec = serde_magnus::deserialize(ruby, hash).map_err(|e| { + Error::new( + runtime_error(), + format!("Failed to deserialize cache config: {e}"), + ) + })?; + regorus::cache::configure(regorus::cache::Config { + regex: spec.regex, + glob: spec.glob, + }); + Ok(()) +} + +#[cfg(feature = "cache")] +fn clear_cache() -> Result<(), Error> { + regorus::cache::clear(); Ok(()) } diff --git a/bindings/ruby/test/test_regorus.rb b/bindings/ruby/test/test_regorus.rb index c06c7cc..b3ea418 100644 --- a/bindings/ruby/test/test_regorus.rb +++ b/bindings/ruby/test/test_regorus.rb @@ -188,6 +188,11 @@ class TestRegorus < Minitest::Test @engine.clear_policy_length_config end + def test_set_cache_config + ::Regorus.set_cache_config({ regex: 256, glob: 128 }) + ::Regorus.clear_cache + end + def alice_results { result: [ diff --git a/bindings/wasm/Cargo.lock b/bindings/wasm/Cargo.lock index c286431..03a092a 100644 --- a/bindings/wasm/Cargo.lock +++ b/bindings/wasm/Cargo.lock @@ -558,6 +558,12 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "lru" +version = "0.16.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1dc47f592c06f33f8e3aea9591776ec7c9f9e4124778ff8a3c3b87159f7e593" + [[package]] name = "memchr" version = "2.7.6" @@ -917,9 +923,11 @@ dependencies = [ "ipnet", "jsonschema", "lazy_static", + "lru", "msvc_spectre_libs", "num-bigint", "num-traits", + "parking_lot", "postcard", "rand", "regex", diff --git a/bindings/wasm/Cargo.toml b/bindings/wasm/Cargo.toml index 42dd89a..edec881 100644 --- a/bindings/wasm/Cargo.toml +++ b/bindings/wasm/Cargo.toml @@ -32,9 +32,11 @@ default = [ "regorus/time", "regorus/uuid", "regorus/urlquery", - "regorus/yaml" + "regorus/yaml", + "cache" ] ast = ["regorus/ast"] +cache = ["regorus/cache"] coverage = ["regorus/coverage"] [dependencies] diff --git a/bindings/wasm/src/lib.rs b/bindings/wasm/src/lib.rs index 6807535..0346017 100644 --- a/bindings/wasm/src/lib.rs +++ b/bindings/wasm/src/lib.rs @@ -35,6 +35,34 @@ struct PolicyLengthSpec { max_lines: usize, } +#[cfg(feature = "cache")] +#[derive(Deserialize)] +struct CacheConfigSpec { + regex: usize, + glob: usize, +} + +/// Configure the global pattern caches used by regex and glob builtins. +/// +/// Accepts a JS object: `{ regex, glob }`. +#[cfg(feature = "cache")] +#[wasm_bindgen(js_name = "setCacheConfig")] +pub fn set_cache_config(config: JsValue) -> Result<(), JsValue> { + let spec: CacheConfigSpec = serde_wasm_bindgen::from_value(config).map_err(error_to_jsvalue)?; + regorus::cache::configure(regorus::cache::Config { + regex: spec.regex, + glob: spec.glob, + }); + Ok(()) +} + +/// Clear all entries from every pattern cache. +#[cfg(feature = "cache")] +#[wasm_bindgen(js_name = "clearCache")] +pub fn clear_cache() { + regorus::cache::clear(); +} + #[wasm_bindgen] pub struct Program { program: Arc, diff --git a/bindings/wasm/test.js b/bindings/wasm/test.js index bbd46b3..04d59bc 100644 --- a/bindings/wasm/test.js +++ b/bindings/wasm/test.js @@ -3,6 +3,9 @@ var regorus = require('./pkg/regorusjs'); +// Configure the global pattern caches. +regorus.setCacheConfig({ regex: 256, glob: 128 }); + // Create an engine. var engine = new regorus.Engine(); diff --git a/src/builtins/glob.rs b/src/builtins/glob.rs index 820e126..a2c2749 100644 --- a/src/builtins/glob.rs +++ b/src/builtins/glob.rs @@ -52,11 +52,33 @@ fn make_delimiters_unix_style(s: &str, delimiters: &[char]) -> Result { } fn make_glob(pattern: &str, span: &Span) -> Result { - Ok(GlobBuilder::new(pattern) - .literal_separator(true) - .build() - .or_else(|_| bail!(span.error("invalid glob")))? - .compile_matcher()) + #[cfg(feature = "cache")] + { + { + let mut cache = crate::cache::GLOB_CACHE.lock(); + if let Some(matcher) = cache.get(pattern) { + return Ok(matcher.clone()); + } + } + let matcher = GlobBuilder::new(pattern) + .literal_separator(true) + .build() + .or_else(|_| bail!(span.error("invalid glob")))? + .compile_matcher(); + { + let mut cache = crate::cache::GLOB_CACHE.lock(); + cache.put(alloc::string::String::from(pattern), matcher.clone()); + } + Ok(matcher) + } + #[cfg(not(feature = "cache"))] + { + Ok(GlobBuilder::new(pattern) + .literal_separator(true) + .build() + .or_else(|_| bail!(span.error("invalid glob")))? + .compile_matcher()) + } } fn glob_match(span: &Span, params: &[Ref], args: &[Value], _strict: bool) -> Result { diff --git a/src/builtins/regex.rs b/src/builtins/regex.rs index 36828f4..46f8d7b 100644 --- a/src/builtins/regex.rs +++ b/src/builtins/regex.rs @@ -12,6 +12,39 @@ use crate::*; use anyhow::{bail, Result}; use regex::Regex; +// --------------------------------------------------------------------------- +// Compiled-regex cache (feature = "cache") +// +// When enabled, compiled Regex objects are stored in a bounded LRU cache +// protected by a Mutex (parking_lot when std, spin when no_std). +// The capacity is configurable at runtime +// via regorus::cache::configure(). +// --------------------------------------------------------------------------- + +/// Compile a regex pattern, using the cache when the `cache` feature +/// is enabled and falling back to direct compilation otherwise. +fn get_or_compile_regex(pattern: &str) -> core::result::Result { + #[cfg(feature = "cache")] + { + { + let mut cache = crate::cache::REGEX_CACHE.lock(); + if let Some(re) = cache.get(pattern) { + return Ok(re.clone()); + } + } + let re = Regex::new(pattern)?; + { + let mut cache = crate::cache::REGEX_CACHE.lock(); + cache.put(alloc::string::String::from(pattern), re.clone()); + Ok(re) + } + } + #[cfg(not(feature = "cache"))] + { + Regex::new(pattern) + } +} + pub fn register(m: &mut builtins::BuiltinsMap<&'static str, builtins::BuiltinFcn>) { m.insert( "regex.find_all_string_submatch_n", @@ -39,8 +72,8 @@ fn find_all_string_submatch_n( let value = ensure_string(name, ¶ms[1], &args[1])?; let n = ensure_numeric(name, ¶ms[2], &args[2])?; - let pattern = - Regex::new(&pattern).or_else(|_| bail!(params[0].span().error("invalid regex")))?; + let re = get_or_compile_regex(&pattern) + .or_else(|_| bail!(params[0].span().error("invalid regex")))?; if !n.is_integer() { bail!(params[2].span().error("n must be an integer")); @@ -53,8 +86,7 @@ fn find_all_string_submatch_n( }; Ok(Value::from_array( - pattern - .captures_iter(&value) + re.captures_iter(&value) .map(|capture| { let groups = capture .iter() @@ -86,8 +118,8 @@ fn find_n(span: &Span, params: &[Ref], args: &[Value], _strict: bool) -> R let value = ensure_string(name, ¶ms[1], &args[1])?; let n = ensure_numeric(name, ¶ms[2], &args[2])?; - let pattern = - Regex::new(&pattern).or_else(|_| bail!(params[0].span().error("invalid regex")))?; + let re = get_or_compile_regex(&pattern) + .or_else(|_| bail!(params[0].span().error("invalid regex")))?; if !n.is_integer() { bail!(params[2].span().error("n must be an integer")); @@ -100,8 +132,7 @@ fn find_n(span: &Span, params: &[Ref], args: &[Value], _strict: bool) -> R }; Ok(Value::from_array( - pattern - .find_iter(&value) + re.find_iter(&value) .map(|m| { let value = Value::String(m.as_str().into()); // Guard match accumulation while pushing each substring. @@ -116,8 +147,11 @@ fn find_n(span: &Span, params: &[Ref], args: &[Value], _strict: bool) -> R fn is_valid(span: &Span, params: &[Ref], args: &[Value], _strict: bool) -> Result { let name = "regex.is_valid"; ensure_args_count(span, name, params, args, 1)?; - Ok(ensure_string(name, ¶ms[0], &args[0]) - .map_or(Value::Bool(false), |p| Value::Bool(Regex::new(&p).is_ok()))) + Ok( + ensure_string(name, ¶ms[0], &args[0]).map_or(Value::Bool(false), |p| { + Value::Bool(get_or_compile_regex(&p).is_ok()) + }), + ) } pub fn regex_match( @@ -131,9 +165,9 @@ pub fn regex_match( let pattern = ensure_string(name, ¶ms[0], &args[0])?; let value = ensure_string(name, ¶ms[1], &args[1])?; - let pattern = - Regex::new(&pattern).or_else(|_| bail!(params[0].span().error("invalid regex")))?; - Ok(Value::Bool(pattern.is_match(&value))) + let re = get_or_compile_regex(&pattern) + .or_else(|_| bail!(params[0].span().error("invalid regex")))?; + Ok(Value::Bool(re.is_match(&value))) } fn regex_replace( @@ -149,15 +183,13 @@ fn regex_replace( let pattern = ensure_string(name, ¶ms[1], &args[1])?; let value = ensure_string(name, ¶ms[2], &args[2])?; - let pattern = match Regex::new(&pattern) { + let re = match get_or_compile_regex(&pattern) { Ok(p) => p, // TODO: This behavior is due to OPA test not raising error. Should we raise error? _ => return Ok(Value::Undefined), }; - Ok(Value::String( - pattern.replace_all(&s, value.as_ref()).into(), - )) + Ok(Value::String(re.replace_all(&s, value.as_ref()).into())) } fn regex_split(span: &Span, params: &[Ref], args: &[Value], _strict: bool) -> Result { @@ -166,11 +198,10 @@ fn regex_split(span: &Span, params: &[Ref], args: &[Value], _strict: bool) let pattern = ensure_string(name, ¶ms[0], &args[0])?; let value = ensure_string(name, ¶ms[1], &args[1])?; - let pattern = - Regex::new(&pattern).or_else(|_| bail!(params[0].span().error("invalid regex")))?; + let re = get_or_compile_regex(&pattern) + .or_else(|_| bail!(params[0].span().error("invalid regex")))?; Ok(Value::from_array( - pattern - .split(&value) + re.split(&value) .map(|s| { let value = Value::String(s.into()); // Guard output accumulation as each split segment is emitted. @@ -211,13 +242,13 @@ fn regex_template_match( } // Fetch pattern, excluding delimiters. - let pattern = Regex::new(&template[start + delimiter_start.len()..end]) + let re = get_or_compile_regex(&template[start + delimiter_start.len()..end]) .or_else(|_| bail!(params[0].span().error("invalid regex")))?; // Skip preceding literal in value. value = &value[start..]; - let m = match pattern.find(value) { + let m = match re.find(value) { Some(m) if m.start() == 0 => m, _ => return Ok(Value::Bool(false)), }; diff --git a/src/cache.rs b/src/cache.rs new file mode 100644 index 0000000..58802d9 --- /dev/null +++ b/src/cache.rs @@ -0,0 +1,171 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Compiled-pattern caches for Rego builtins. +//! +//! When the `cache` feature is enabled, compiled [`regex::Regex`] and +//! [`globset::GlobMatcher`] objects are held in bounded LRU caches so that +//! repeated evaluations of the same pattern avoid recompilation. +//! +//! # Examples +//! +//! ```ignore +//! use regorus::cache; +//! +//! // Configure cache capacities (0 = disabled). +//! cache::configure(cache::Config { +//! regex: 256, +//! glob: 128, +//! }); +//! +//! // Flush all cached patterns. +//! cache::clear(); +//! ``` + +#[cfg(any(feature = "regex", feature = "glob"))] +use core::num::NonZeroUsize; +#[cfg(any(feature = "regex", feature = "glob"))] +use lazy_static::lazy_static; +#[cfg(all(feature = "std", any(feature = "regex", feature = "glob")))] +use parking_lot::Mutex; +#[cfg(all(not(feature = "std"), any(feature = "regex", feature = "glob")))] +use spin::Mutex; + +#[cfg(any(feature = "regex", feature = "glob"))] +use alloc::string::String; + +/// Configuration for builtin pattern caches. +/// +/// Each field controls the maximum number of compiled patterns held in the +/// corresponding LRU cache. A value of `0` disables that cache entirely +/// (every lookup recompiles). Values exceeding [`Config::MAX_CAPACITY`] are +/// clamped silently. +#[derive(Debug, Clone, Copy)] +pub struct Config { + /// Maximum compiled regex patterns (default 256). + pub regex: usize, + /// Maximum compiled glob matchers (default 128). + pub glob: usize, +} + +impl Config { + /// Hard upper bound for any single cache capacity (2^16 = 65 536). + pub const MAX_CAPACITY: usize = 1 << 16; +} + +impl Default for Config { + fn default() -> Self { + Self { + regex: 256, + glob: 128, + } + } +} + +// --------------------------------------------------------------------------- +// Internal generic LRU wrapper +// --------------------------------------------------------------------------- + +#[cfg(any(feature = "regex", feature = "glob"))] +pub(crate) struct LruCache { + inner: Option>, +} + +#[cfg(any(feature = "regex", feature = "glob"))] +impl LruCache { + pub(crate) fn new(capacity: usize) -> Self { + Self { + inner: NonZeroUsize::new(capacity).map(lru::LruCache::new), + } + } + + /// Look up a key, returning a reference if present. Promotes to most-recent. + pub(crate) fn get(&mut self, key: &str) -> Option<&V> { + self.inner.as_mut()?.get(key) + } + + /// Insert a key-value pair. Evicts the least-recently-used entry if full. + pub(crate) fn put(&mut self, key: String, value: V) { + if let Some(cache) = self.inner.as_mut() { + cache.put(key, value); + } + } + + /// Remove all entries. + pub(crate) fn clear(&mut self) { + if let Some(cache) = self.inner.as_mut() { + cache.clear(); + } + } + + /// Resize the cache. If new capacity is 0, disables the cache. + pub(crate) fn resize(&mut self, capacity: usize) { + match NonZeroUsize::new(capacity) { + Some(cap) => match self.inner.as_mut() { + Some(cache) => cache.resize(cap), + None => self.inner = Some(lru::LruCache::new(cap)), + }, + None => { + self.inner = None; + } + } + } + + /// Number of entries currently cached. + #[allow(dead_code)] + pub(crate) fn len(&self) -> usize { + self.inner.as_ref().map_or(0, lru::LruCache::len) + } +} + +// --------------------------------------------------------------------------- +// Global regex cache +// --------------------------------------------------------------------------- + +#[cfg(feature = "regex")] +lazy_static! { + pub(crate) static ref REGEX_CACHE: Mutex> = + Mutex::new(LruCache::new(Config::default().regex)); +} + +// --------------------------------------------------------------------------- +// Global glob cache +// --------------------------------------------------------------------------- + +#[cfg(feature = "glob")] +lazy_static! { + pub(crate) static ref GLOB_CACHE: Mutex> = + Mutex::new(LruCache::new(Config::default().glob)); +} + +// --------------------------------------------------------------------------- +// Public API +// --------------------------------------------------------------------------- + +/// Apply a new cache configuration. +/// +/// Resizes each cache to the specified capacity. Existing entries are +/// preserved (subject to LRU eviction if the new capacity is smaller). +/// Values exceeding [`Config::MAX_CAPACITY`] are clamped. +pub fn configure(config: Config) { + let regex = config.regex.min(Config::MAX_CAPACITY); + let glob = config.glob.min(Config::MAX_CAPACITY); + + #[cfg(feature = "regex")] + REGEX_CACHE.lock().resize(regex); + + #[cfg(feature = "glob")] + GLOB_CACHE.lock().resize(glob); + + // Suppress unused-variable warnings when neither regex nor glob is enabled. + let _ = (regex, glob); +} + +/// Remove all entries from every pattern cache. +pub fn clear() { + #[cfg(feature = "regex")] + REGEX_CACHE.lock().clear(); + + #[cfg(feature = "glob")] + GLOB_CACHE.lock().clear(); +} diff --git a/src/interpreter.rs b/src/interpreter.rs index 6f90382..04b0eb1 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -289,7 +289,7 @@ impl Interpreter { } fn execution_timer_tick(&mut self, work_units: u32) -> Result<()> { - if self.execution_timer.limit().is_none() { + if !self.execution_timer.accumulate(work_units) { return Ok(()); } @@ -297,7 +297,7 @@ impl Interpreter { return Ok(()); }; - self.execution_timer.tick(work_units, now)?; + self.execution_timer.check_now(now)?; Ok(()) } diff --git a/src/languages/rego/compiler/destructuring.rs b/src/languages/rego/compiler/destructuring.rs index 674878c..69cc298 100644 --- a/src/languages/rego/compiler/destructuring.rs +++ b/src/languages/rego/compiler/destructuring.rs @@ -91,6 +91,19 @@ impl<'a> Compiler<'a> { AssignmentPlan::EqualityCheck { lhs_expr, rhs_expr } => { let lhs_reg = self.compile_rego_expr_with_span(lhs_expr, lhs_expr.span(), false)?; let rhs_reg = self.compile_rego_expr_with_span(rhs_expr, rhs_expr.span(), false)?; + if !self.soft_assert_mode { + // AssertEq handles the equality assertion inline; the returned + // register is not used as a boolean by the caller — it is the + // expression's "result register" for potential downstream use. + self.emit_instruction( + Instruction::AssertEq { + left: lhs_reg, + right: rhs_reg, + }, + span, + ); + return Ok(lhs_reg); + } let dest = self.alloc_register(); self.emit_instruction( Instruction::Eq { @@ -100,9 +113,6 @@ impl<'a> Compiler<'a> { }, span, ); - if !self.soft_assert_mode { - self.emit_instruction(Instruction::AssertCondition { condition: dest }, span); - } Ok(dest) } AssignmentPlan::WildcardMatch { @@ -196,35 +206,47 @@ impl<'a> Compiler<'a> { DestructuringPlan::EqualityExpr(expected_expr) => { let expected_reg = self.compile_rego_expr_with_span(expected_expr, expected_expr.span(), false)?; - let cmp_reg = self.alloc_register(); + if self.soft_assert_mode { + let cmp_reg = self.alloc_register(); + self.emit_instruction( + Instruction::Eq { + dest: cmp_reg, + left: value_register, + right: expected_reg, + }, + span, + ); + return Ok(Some(cmp_reg)); + } self.emit_instruction( - Instruction::Eq { - dest: cmp_reg, + Instruction::AssertEq { left: value_register, right: expected_reg, }, span, ); - if self.soft_assert_mode { - return Ok(Some(cmp_reg)); - } - self.emit_instruction(Instruction::AssertCondition { condition: cmp_reg }, span); } DestructuringPlan::EqualityValue(expected_value) => { let expected_reg = self.load_literal_value(expected_value, span); - let cmp_reg = self.alloc_register(); + if self.soft_assert_mode { + let cmp_reg = self.alloc_register(); + self.emit_instruction( + Instruction::Eq { + dest: cmp_reg, + left: value_register, + right: expected_reg, + }, + span, + ); + return Ok(Some(cmp_reg)); + } self.emit_instruction( - Instruction::Eq { - dest: cmp_reg, + Instruction::AssertEq { left: value_register, right: expected_reg, }, span, ); - if self.soft_assert_mode { - return Ok(Some(cmp_reg)); - } - self.emit_instruction(Instruction::AssertCondition { condition: cmp_reg }, span); } DestructuringPlan::Array { element_plans } => { self.assert_array_length(value_register, element_plans.len(), span)?; @@ -371,16 +393,13 @@ impl<'a> Compiler<'a> { span, ); - let cmp_reg = self.alloc_register(); self.emit_instruction( - Instruction::Eq { - dest: cmp_reg, + Instruction::AssertEq { left: actual_len_reg, right: expected_len_reg, }, span, ); - self.emit_instruction(Instruction::AssertCondition { condition: cmp_reg }, span); Ok(()) } } diff --git a/src/languages/rego/compiler/expressions.rs b/src/languages/rego/compiler/expressions.rs index 3f39214..ff184a6 100644 --- a/src/languages/rego/compiler/expressions.rs +++ b/src/languages/rego/compiler/expressions.rs @@ -5,6 +5,8 @@ mod collection_literals; mod operations; +pub(super) use collection_literals::try_eval_const; + use super::{Compiler, CompilerError, Register, Result}; use crate::ast::{Expr, ExprRef}; use crate::compiler::destructuring_planner::plans::BindingPlan; diff --git a/src/languages/rego/compiler/expressions/collection_literals.rs b/src/languages/rego/compiler/expressions/collection_literals.rs index 1ff8871..ac9ea91 100644 --- a/src/languages/rego/compiler/expressions/collection_literals.rs +++ b/src/languages/rego/compiler/expressions/collection_literals.rs @@ -7,20 +7,62 @@ )] use super::{Compiler, Register, Result}; -use crate::ast::ExprRef; +use crate::ast::{Expr, ExprRef}; use crate::lexer::Span; use crate::rvm::instructions::{ArrayCreateParams, ObjectCreateParams, SetCreateParams}; use crate::rvm::Instruction; use crate::{Rc, Value}; -use alloc::collections::BTreeMap; +use alloc::collections::{BTreeMap, BTreeSet}; use alloc::vec::Vec; +/// Try to evaluate an expression as a compile-time constant. +pub(in crate::languages::rego::compiler) fn try_eval_const(expr: &Expr) -> Option { + match expr { + Expr::Number { value, .. } + | Expr::String { value, .. } + | Expr::RawString { value, .. } + | Expr::Bool { value, .. } + | Expr::Null { value, .. } => Some(value.clone()), + Expr::UnaryExpr { expr, .. } => match expr.as_ref() { + Expr::Number { + value: Value::Number(n), + .. + } => Some(Value::Number(n.neg()?)), + _ => None, + }, + Expr::Array { items, .. } => items + .iter() + .map(|i| try_eval_const(i.as_ref())) + .collect::>>() + .map(|v| Value::Array(Rc::new(v))), + Expr::Set { items, .. } => items + .iter() + .map(|i| try_eval_const(i.as_ref())) + .collect::>>() + .map(|s| Value::Set(Rc::new(s))), + Expr::Object { fields, .. } => fields + .iter() + .map(|(_, k, v)| Some((try_eval_const(k.as_ref())?, try_eval_const(v.as_ref())?))) + .collect::>>() + .map(|m| Value::Object(Rc::new(m))), + _ => None, + } +} + impl<'a> Compiler<'a> { pub(super) fn compile_array_literal( &mut self, items: &[ExprRef], span: &Span, ) -> Result { + let all_const: Option> = items.iter().map(|i| try_eval_const(i.as_ref())).collect(); + if let Some(values) = all_const { + let dest = self.alloc_register(); + let literal_idx = self.add_literal(Value::Array(Rc::new(values))); + self.emit_instruction(Instruction::Load { dest, literal_idx }, span); + return Ok(dest); + } + let mut element_registers = Vec::with_capacity(items.len()); for item in items { let item_reg = self.compile_rego_expr_with_span(item, item.span(), false)?; @@ -45,6 +87,15 @@ impl<'a> Compiler<'a> { items: &[ExprRef], span: &Span, ) -> Result { + let all_const: Option> = + items.iter().map(|i| try_eval_const(i.as_ref())).collect(); + if let Some(values) = all_const { + let dest = self.alloc_register(); + let literal_idx = self.add_literal(Value::Set(Rc::new(values))); + self.emit_instruction(Instruction::Load { dest, literal_idx }, span); + return Ok(dest); + } + let mut element_registers = Vec::with_capacity(items.len()); for item in items { let item_reg = self.compile_rego_expr_with_span(item, item.span(), false)?; @@ -66,6 +117,17 @@ impl<'a> Compiler<'a> { fields: &[(crate::lexer::Span, ExprRef, ExprRef)], span: &Span, ) -> Result { + let all_const: Option> = fields + .iter() + .map(|(_, k, v)| Some((try_eval_const(k.as_ref())?, try_eval_const(v.as_ref())?))) + .collect(); + if let Some(obj) = all_const { + let dest = self.alloc_register(); + let literal_idx = self.add_literal(Value::Object(Rc::new(obj))); + self.emit_instruction(Instruction::Load { dest, literal_idx }, span); + return Ok(dest); + } + let dest = self.alloc_register(); let mut value_regs = Vec::with_capacity(fields.len()); diff --git a/src/languages/rego/compiler/mod.rs b/src/languages/rego/compiler/mod.rs index 3faea56..c54c28a 100644 --- a/src/languages/rego/compiler/mod.rs +++ b/src/languages/rego/compiler/mod.rs @@ -23,6 +23,7 @@ pub use error::{CompilerError, Result, SpannedCompilerError}; use crate::ast::ExprRef; use crate::lexer::Span; use crate::rvm::program::{Program, RuleType, SpanInfo}; +use crate::value::Value; use crate::CompiledPolicy; use alloc::collections::{BTreeMap, BTreeSet}; use alloc::string::String; @@ -120,6 +121,10 @@ pub struct Compiler<'a> { rule_definitions: Vec>>, rule_definition_function_params: Vec>>>, rule_definition_destructuring_patterns: Vec>>, + /// Per-rule, per-definition: the static value produced by this definition, + /// or `None` if the value is dynamic or differs across else-branches. + /// Used to compute `RuleInfo::early_exit_on_first_success`. + rule_definition_static_values: Vec>>, rule_types: Vec, rule_function_param_count: Vec>, rule_result_registers: Vec, @@ -153,6 +158,7 @@ impl<'a> Compiler<'a> { rule_definitions: Vec::new(), rule_definition_function_params: Vec::new(), rule_definition_destructuring_patterns: Vec::new(), + rule_definition_static_values: Vec::new(), rule_types: Vec::new(), rule_function_param_count: Vec::new(), rule_result_registers: Vec::new(), diff --git a/src/languages/rego/compiler/program.rs b/src/languages/rego/compiler/program.rs index 350b519..ff39dce 100644 --- a/src/languages/rego/compiler/program.rs +++ b/src/languages/rego/compiler/program.rs @@ -99,6 +99,38 @@ impl<'a> Compiler<'a> { }; rule_info.destructuring_blocks = destructuring_blocks; + + // Compute early_exit_on_first_success: if every definition has + // the same static value, the VM can stop after the first success. + // Only relevant for Complete rules and function rules with ≥2 defs. + let static_values = self.rule_definition_static_values.get(rule_index as usize); + if let Some(svs) = static_values { + if svs.len() >= 2 { + // Check that all definitions have a known static value and + // that they are all equal. + let mut all_same = true; + let mut reference: Option<&Value> = None; + for sv in svs { + match sv { + Some(v) => match reference { + None => reference = Some(v), + Some(r) => { + if r != v { + all_same = false; + break; + } + } + }, + None => { + all_same = false; + break; + } + } + } + rule_info.early_exit_on_first_success = all_same && reference.is_some(); + } + } + rule_infos_map.insert(rule_index as usize, rule_info); } diff --git a/src/languages/rego/compiler/queries.rs b/src/languages/rego/compiler/queries.rs index f781b67..4043a84 100644 --- a/src/languages/rego/compiler/queries.rs +++ b/src/languages/rego/compiler/queries.rs @@ -242,21 +242,7 @@ impl<'a> Compiler<'a> { compiler.compile_rego_expr_with_span(expr, expr.span(), false) })?; - let negated_reg = self.alloc_register(); - self.emit_instruction( - Instruction::Not { - dest: negated_reg, - operand: expr_reg, - }, - &stmt.span, - ); - - self.emit_instruction( - Instruction::AssertCondition { - condition: negated_reg, - }, - &stmt.span, - ); + self.emit_instruction(Instruction::AssertNot { operand: expr_reg }, &stmt.span); } } Ok(()) diff --git a/src/languages/rego/compiler/rules.rs b/src/languages/rego/compiler/rules.rs index dacdc99..6fa12f9 100644 --- a/src/languages/rego/compiler/rules.rs +++ b/src/languages/rego/compiler/rules.rs @@ -11,7 +11,7 @@ )] use super::{CompilationContext, Compiler, CompilerError, ContextType, Result, WorklistEntry}; -use crate::ast::{Expr, Rule, RuleHead}; +use crate::ast::{Expr, ExprRef, Rule, RuleHead}; use crate::compiler::destructuring_planner::plans::BindingPlan; use crate::lexer::Span; use crate::rvm::program::{Program, RuleType}; @@ -26,6 +26,16 @@ use alloc::sync::Arc; use alloc::vec::Vec; impl<'a> Compiler<'a> { + /// Extract a compile-time constant `Value` from an optional expression. + /// Returns `Some(Value::Bool(true))` for the implicit-true case (`expr_ref` + /// is `None`), delegates to `try_eval_const` for actual expressions. + fn static_value_of_expr(expr_ref: &Option) -> Option { + match expr_ref { + None => Some(Value::Bool(true)), + Some(expr) => super::expressions::try_eval_const(expr.as_ref()), + } + } + pub(super) fn compute_rule_type(&self, rule_path: &str) -> Result { let Some(definitions) = self.policy.inner.rules.get(rule_path) else { return Err(CompilerError::General { @@ -311,6 +321,10 @@ impl<'a> Compiler<'a> { self.rule_definition_destructuring_patterns.push(Vec::new()); } + while self.rule_definition_static_values.len() <= rule_index as usize { + self.rule_definition_static_values.push(Vec::new()); + } + let mut num_registers_used = 0; let mut rule_param_count: Option = None; @@ -525,6 +539,54 @@ impl<'a> Compiler<'a> { self.pop_scope(); + // Compute this definition's static value for early-exit analysis. + // A definition has a known static value if every body (including + // else-branches) would produce the same literal. + let def_static_value = if bodies.is_empty() { + // No bodies — value comes from the head's value_expr. + let head_value = self + .context_stack + .last() + .and_then(|ctx| ctx.value_expr.clone()); + Self::static_value_of_expr(&head_value) + } else { + // Replay the same value_expr resolution as the body loop. + let head_value = self + .context_stack + .last() + .and_then(|ctx| ctx.value_expr.clone()); + let mut consistent: Option = None; + let mut all_same = true; + for (bi, b) in bodies.iter().enumerate() { + let mut bve: Option = + b.assign.as_ref().map(|a| a.value.clone()); + if bve.is_none() && bi == 0 { + bve = head_value.clone(); + } + match Self::static_value_of_expr(&bve) { + Some(v) => match &consistent { + None => consistent = Some(v), + Some(prev) => { + if *prev != v { + all_same = false; + break; + } + } + }, + None => { + all_same = false; + break; + } + } + } + if all_same { + consistent + } else { + None + } + }; + self.rule_definition_static_values[rule_index as usize].push(def_static_value); + self.rule_definitions[rule_index as usize].push(body_entry_points); if self.register_counter > num_registers_used { diff --git a/src/lib.rs b/src/lib.rs index 2c45b60..8f23add 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -176,6 +176,25 @@ pub use utils::limits::{ }; pub use value::Value; +/// Compiled-pattern caches for the `regex.*` and `glob.*` Rego builtins. +/// +/// When the `cache` feature is enabled, compiled patterns are held in +/// bounded LRU caches so that repeated evaluations avoid recompilation. +/// +/// # Examples +/// +/// ```ignore +/// use regorus::cache; +/// +/// cache::configure(cache::Config { +/// regex: 256, +/// glob: 128, +/// }); +/// cache::clear(); +/// ``` +#[cfg(feature = "cache")] +pub mod cache; + #[cfg(feature = "arc")] pub use alloc::sync::Arc as Rc; diff --git a/src/rvm/instructions/display.rs b/src/rvm/instructions/display.rs index b681cba..73ee97e 100644 --- a/src/rvm/instructions/display.rs +++ b/src/rvm/instructions/display.rs @@ -242,6 +242,12 @@ impl core::fmt::Display for Instruction { Instruction::Count { dest, collection } => { format!("COUNT R({}) R({})", dest, collection) } + Instruction::AssertEq { left, right } => { + format!("ASSERT_EQ R({}) R({})", left, right) + } + Instruction::AssertNot { operand } => { + format!("ASSERT_NOT R({})", operand) + } Instruction::AssertCondition { condition } => { format!("ASSERT_CONDITION R({})", condition) } diff --git a/src/rvm/instructions/mod.rs b/src/rvm/instructions/mod.rs index d829038..ea8af94 100644 --- a/src/rvm/instructions/mod.rs +++ b/src/rvm/instructions/mod.rs @@ -131,6 +131,8 @@ pub enum Instruction { left: u8, right: u8, }, + /// Rego negation - produces `true` if operand is `false` or undefined, + /// `false` for any other defined value (including non-booleans). Not { dest: u8, operand: u8, @@ -243,6 +245,17 @@ pub enum Instruction { collection: u8, }, + /// Assert that two registers are equal - if either is undefined or they differ, fail the condition + AssertEq { + left: u8, + right: u8, + }, + + /// Assert negation - succeed if operand is false or undefined, fail if true + AssertNot { + operand: u8, + }, + /// Assert condition - if register contains false or undefined, return undefined immediately AssertCondition { condition: u8, diff --git a/src/rvm/program/core.rs b/src/rvm/program/core.rs index 05c781f..4ee66fa 100644 --- a/src/rvm/program/core.rs +++ b/src/rvm/program/core.rs @@ -85,7 +85,7 @@ pub struct Program { impl Program { /// Current serialization format version - pub const SERIALIZATION_VERSION: u32 = 4; + pub const SERIALIZATION_VERSION: u32 = 5; /// Magic bytes to identify Regorus program files pub const MAGIC: [u8; 4] = *b"REGO"; /// Maximum instructions supported (matches u16 jump targets) diff --git a/src/rvm/program/listing.rs b/src/rvm/program/listing.rs index 5c19f4f..a79d7f6 100644 --- a/src/rvm/program/listing.rs +++ b/src/rvm/program/listing.rs @@ -358,7 +358,10 @@ fn format_instruction_readable( } Instruction::Not { dest, operand } => { let base = format!("{}Not r{} ← ¬r{}", indent, dest, operand); - let comment = format!("Logical NOT: !r{}", operand); + let comment = format!( + "Rego negation: true if r{} is false/undefined, false otherwise", + operand + ); align_comment(&base, &comment, config.comment_column) } Instruction::BuiltinCall { params_index } => { @@ -575,6 +578,22 @@ fn format_instruction_readable( let comment = format!("Get count/length of collection r{}", collection); align_comment(&base, &comment, config.comment_column) } + Instruction::AssertEq { left, right } => { + let base = format!("{}AssertEq assert r{} == r{}", indent, left, right); + let comment = format!( + "Assert r{} equals r{} (exit if unequal/undefined)", + left, right + ); + align_comment(&base, &comment, config.comment_column) + } + Instruction::AssertNot { operand } => { + let base = format!("{}AssertNot assert !r{}", indent, operand); + let comment = format!( + "Assert r{} is false/undefined (exit if any defined truthy value)", + operand + ); + align_comment(&base, &comment, config.comment_column) + } Instruction::AssertCondition { condition } => { let base = format!("{}Assert assert r{}", indent, condition); let comment = format!("Assert r{} is true (exit if false/undefined)", condition); @@ -899,6 +918,8 @@ const fn get_instruction_name(instruction: &Instruction) -> &'static str { Instruction::SetCreate { .. } => "SET_CREATE", Instruction::Contains { .. } => "CONTAINS", Instruction::Count { .. } => "COUNT", + Instruction::AssertEq { .. } => "ASSERT_EQ", + Instruction::AssertNot { .. } => "ASSERT_NOT", Instruction::AssertCondition { .. } => "ASSERT", Instruction::AssertNotUndefined { .. } => "ASSERT_NOT_UNDEF", Instruction::LoopStart { .. } => "LOOP_START", diff --git a/src/rvm/program/serialization/binary.rs b/src/rvm/program/serialization/binary.rs index 85f2976..f5fad82 100644 --- a/src/rvm/program/serialization/binary.rs +++ b/src/rvm/program/serialization/binary.rs @@ -157,7 +157,7 @@ impl Program { program.rego_v0 = Self::legacy_rego_v0(data, version).unwrap_or(false); Ok(DeserializationResult::Partial(program)) } - 4 => { + 4 | 5 => { if data.len() < 29 { return Err("Data too short for header".to_string()); } @@ -281,7 +281,7 @@ impl Program { let version = Self::read_u32(data, 4).ok(); match version { - Some(1..=4) => Ok(true), + Some(1..=5) => Ok(true), _ => Ok(false), } } diff --git a/src/rvm/program/types.rs b/src/rvm/program/types.rs index 72869d9..a2a78fb 100644 --- a/src/rvm/program/types.rs +++ b/src/rvm/program/types.rs @@ -97,6 +97,11 @@ pub struct RuleInfo { /// Optional destructuring block entry point per definition /// Index: definition_index → Some(entry_point) | None pub destructuring_blocks: Vec>, + /// If true, all definitions are statically known to produce the same result + /// value, so the VM can stop after the first successful definition without + /// checking consistency with remaining definitions. + #[serde(default)] + pub early_exit_on_first_success: bool, } impl RuleInfo { @@ -117,6 +122,7 @@ impl RuleInfo { result_reg, num_registers, destructuring_blocks: alloc::vec![None; num_definitions], + early_exit_on_first_success: false, } } @@ -143,6 +149,7 @@ impl RuleInfo { result_reg, num_registers, destructuring_blocks: alloc::vec![None; num_definitions], + early_exit_on_first_success: false, } } diff --git a/src/rvm/tests/vm.rs b/src/rvm/tests/vm.rs index f293985..c899071 100644 --- a/src/rvm/tests/vm.rs +++ b/src/rvm/tests/vm.rs @@ -414,6 +414,7 @@ mod tests { result_reg, num_registers: 50, // Increased to accommodate test cases with higher register indices destructuring_blocks, + early_exit_on_first_success: false, }; program.rule_infos.push(rule_info); diff --git a/src/rvm/vm/dispatch.rs b/src/rvm/vm/dispatch.rs index 02b9313..78d96c4 100644 --- a/src/rvm/vm/dispatch.rs +++ b/src/rvm/vm/dispatch.rs @@ -26,7 +26,6 @@ impl RegoVM { program: &Program, instruction: Instruction, ) -> Result { - self.memory_check()?; self.execute_load_and_move(program, instruction) } @@ -319,25 +318,32 @@ impl RegoVM { Not { dest, operand } => { let operand_value = self.get_register(operand)?; - if operand_value == &Value::Undefined { - // In Rego, `not expr` succeeds when `expr` has no results. - // When the operand evaluates to undefined we should treat it as - // a successful negation instead of propagating undefined. - self.set_register(dest, Value::Bool(true))?; - return Ok(InstructionOutcome::Continue); - } - - if let Some(value) = self.to_bool(operand_value) { - self.set_register(dest, Value::Bool(!value))?; - Ok(InstructionOutcome::Continue) - } else { - Err(VmError::ArithmeticError { - message: alloc::format!( - "#undefined: logical NOT expects a boolean (operand={operand_value:?})" - ), - pc: self.pc, - }) - } + // In Rego, `not expr` succeeds when `expr` is undefined or false, + // and fails for any other defined value (including non-booleans like 42). + let negated = match *operand_value { + Value::Undefined => true, + Value::Bool(b) => !b, + _ => false, + }; + self.set_register(dest, Value::Bool(negated))?; + Ok(InstructionOutcome::Continue) + } + AssertEq { left, right } => { + let a = self.get_register(left)?; + let b = self.get_register(right)?; + let passed = a != &Value::Undefined && b != &Value::Undefined && a == b; + self.handle_condition(passed)?; + Ok(InstructionOutcome::Continue) + } + AssertNot { operand } => { + let value = self.get_register(operand)?; + let passed = match *value { + Value::Undefined => true, + Value::Bool(b) => !b, + _ => false, + }; + self.handle_condition(passed)?; + Ok(InstructionOutcome::Continue) } AssertCondition { condition } => { let value = self.get_register(condition)?; @@ -696,10 +702,9 @@ impl RegoVM { Value::Array(ref array_items) => { Value::Bool(array_items.contains(value_to_check)) } - Value::Object(ref object_fields) => Value::Bool( - object_fields.contains_key(value_to_check) - || object_fields.values().any(|v| v == value_to_check), - ), + Value::Object(ref object_fields) => { + Value::Bool(object_fields.values().any(|v| v == value_to_check)) + } _ => Value::Bool(false), }; diff --git a/src/rvm/vm/execution.rs b/src/rvm/vm/execution.rs index 8400714..8b2a049 100644 --- a/src/rvm/vm/execution.rs +++ b/src/rvm/vm/execution.rs @@ -161,6 +161,7 @@ impl RegoVM { self.reset_execution_state(); self.reset_execution_timer_state(); self.execution_state = ExecutionState::Running; + self.enforce_memory_check()?; match self.jump_to(0_u32) { Ok(value) => { self.execution_state = ExecutionState::Completed { @@ -179,6 +180,7 @@ impl RegoVM { self.reset_execution_state(); self.reset_execution_timer_state(); self.execution_state = ExecutionState::Running; + self.enforce_memory_check()?; match self.run_stackless_from(0) { Ok(result) => Ok(result), Err(err) => { @@ -191,6 +193,7 @@ impl RegoVM { fn execute_suspendable_entry(&mut self, entry_point_pc: usize) -> Result { self.execution_state = ExecutionState::Running; self.reset_execution_timer_state(); + self.enforce_memory_check()?; match self.run_stackless_from(entry_point_pc) { Ok(result) => Ok(result), Err(err) => { diff --git a/src/rvm/vm/machine.rs b/src/rvm/vm/machine.rs index 9d931be..6baf6d2 100644 --- a/src/rvm/vm/machine.rs +++ b/src/rvm/vm/machine.rs @@ -390,7 +390,7 @@ impl RegoVM { } pub(super) fn execution_timer_tick(&mut self, work_units: u32) -> Result<()> { - if self.execution_timer.limit().is_none() { + if !self.execution_timer.accumulate(work_units) { return Ok(()); } @@ -399,7 +399,7 @@ impl RegoVM { }; self.execution_timer - .tick(work_units, now) + .check_now(now) .map_err(|err| match err { LimitError::TimeLimitExceeded { elapsed, limit } => VmError::TimeLimitExceeded { elapsed, @@ -505,8 +505,8 @@ impl RegoVM { } #[cfg(all(feature = "allocator-memory-limits", not(miri)))] - pub(super) fn memory_check(&mut self) -> Result<()> { - limits::check_memory_limit_if_needed().map_err(|err| match err { + fn map_limit_error(&self, err: LimitError) -> VmError { + match err { LimitError::MemoryLimitExceeded { usage, limit } => VmError::MemoryLimitExceeded { usage, limit, @@ -516,7 +516,17 @@ impl RegoVM { message: format!("unexpected limit error: {other}"), pc: self.pc, }, - }) + } + } + + #[cfg(all(feature = "allocator-memory-limits", not(miri)))] + pub(super) fn memory_check(&mut self) -> Result<()> { + limits::check_memory_limit_if_needed().map_err(|err| self.map_limit_error(err)) + } + + #[cfg(all(feature = "allocator-memory-limits", not(miri)))] + pub(super) fn enforce_memory_check(&mut self) -> Result<()> { + limits::enforce_memory_limit().map_err(|err| self.map_limit_error(err)) } #[cfg(any(miri, not(feature = "allocator-memory-limits")))] @@ -524,6 +534,11 @@ impl RegoVM { Ok(()) } + #[cfg(any(miri, not(feature = "allocator-memory-limits")))] + pub(super) fn enforce_memory_check(&mut self) -> Result<()> { + Ok(()) + } + /// Get or create the cached dummy span for builtin calls. pub(super) fn get_dummy_span(&mut self) -> Result<&crate::lexer::Span> { if self.dummy_span.is_none() { diff --git a/src/rvm/vm/rules.rs b/src/rvm/vm/rules.rs index d9388bc..4276609 100644 --- a/src/rvm/vm/rules.rs +++ b/src/rvm/vm/rules.rs @@ -98,6 +98,11 @@ impl RegoVM { } } else { first_successful_result = Some(current_result.clone()); + // All definitions produce the same static value; + // no need to verify consistency with the rest. + if rule_info.early_exit_on_first_success { + break 'outer; + } } } } @@ -607,6 +612,13 @@ impl RegoVM { } } else { frame_data.accumulated_result = Some(current_result); + // All definitions produce the same static value; + // skip remaining definitions. + if rule_info.early_exit_on_first_success { + frame_data.current_definition_index = frame_data.total_definitions; + frame_data.phase = RuleFramePhase::Finalizing; + return Ok(None); + } } } } diff --git a/src/utils/limits/time.rs b/src/utils/limits/time.rs index 36b9055..94be970 100644 --- a/src/utils/limits/time.rs +++ b/src/utils/limits/time.rs @@ -247,6 +247,26 @@ impl ExecutionTimer { self.last_elapsed } + /// Increment work units and return whether a time check is due. + /// + /// This method only updates the internal counter — it never reads a clock. + /// Callers should obtain the current time and call [`check_now`](Self::check_now) + /// only when this returns `true`. + pub const fn accumulate(&mut self, work_units: u32) -> bool { + let Some(config) = self.config else { + return false; + }; + self.accumulated_units = self.accumulated_units.saturating_add(work_units); + if self.accumulated_units < config.check_interval.get() { + return false; + } + + // Preserve the remainder so that callers do not lose fractional work. + let interval = config.check_interval.get(); + self.accumulated_units %= interval; + true + } + /// Increment work units and run the periodic limit check when necessary. pub fn tick(&mut self, work_units: u32, now: Duration) -> Result<(), LimitError> { let Some(config) = self.config else { @@ -471,4 +491,50 @@ mod tests { let mut slot = super::TIME_SOURCE_OVERRIDE.lock(); *slot = previous; } + + #[test] + fn accumulate_defers_clock_reads() { + let mut timer = ExecutionTimer::new(Some(ExecutionTimerConfig { + limit: Duration::from_secs(1), + check_interval: nz(4), + })); + timer.start(Duration::from_millis(0)); + + // First 3 work units should not require a clock read. + for _ in 0..3 { + assert!( + !timer.accumulate(1), + "accumulate before interval must return false" + ); + } + + // The 4th unit crosses the interval — caller should read the clock now. + assert!( + timer.accumulate(1), + "accumulate at interval must return true" + ); + + // After the boundary, the counter resets — next 3 units are cheap again. + for _ in 0..3 { + assert!( + !timer.accumulate(1), + "accumulate after reset must return false" + ); + } + assert!( + timer.accumulate(1), + "second interval crossing must return true" + ); + } + + #[test] + fn accumulate_returns_false_when_disabled() { + let mut timer = ExecutionTimer::new(None); + for _ in 0..10 { + assert!( + !timer.accumulate(1), + "disabled timer must never request a clock read" + ); + } + } } diff --git a/tests/rvm/compiler.rs b/tests/rvm/compiler.rs new file mode 100644 index 0000000..b490cb3 --- /dev/null +++ b/tests/rvm/compiler.rs @@ -0,0 +1,343 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. +#![cfg(feature = "rvm")] + +use regorus::languages::rego::compiler::Compiler; +use regorus::rvm::Instruction; +use regorus::{Engine, Rc, Value}; +use std::collections::BTreeSet; + +/// Compile a single-rule Rego module and return the program. +fn compile_rule(module: &str) -> std::sync::Arc { + let mut engine = Engine::new(); + engine + .add_policy("test.rego".to_string(), module.to_string()) + .expect("failed to add policy"); + let compiled = engine + .compile_with_entrypoint(&Rc::from("data.test.p")) + .expect("failed to compile policy"); + Compiler::compile_from_policy(&compiled, &["data.test.p"]).expect("failed to compile to RVM") +} + +/// Assert that the program's instruction stream contains no collection-create +/// instructions (ArrayCreate, SetCreate, ObjectCreate), meaning the collections +/// were hoisted into the literal table. +fn assert_no_collection_create(program: ®orus::rvm::program::Program) { + for (pc, instr) in program.instructions.iter().enumerate() { + match instr { + Instruction::ArrayCreate { .. } => { + panic!("unexpected ArrayCreate at pc={pc}; expected hoisted constant") + } + Instruction::SetCreate { .. } => { + panic!("unexpected SetCreate at pc={pc}; expected hoisted constant") + } + Instruction::ObjectCreate { .. } => { + panic!("unexpected ObjectCreate at pc={pc}; expected hoisted constant") + } + _ => {} + } + } +} + +/// Assert the literal table contains a value equal to `expected`. +fn assert_literal_exists(program: ®orus::rvm::program::Program, expected: &Value) { + assert!( + program.literals.iter().any(|v| v == expected), + "expected literal {:?} not found in literal table: {:?}", + expected, + program.literals + ); +} + +#[test] +fn constant_array_is_hoisted() { + let program = compile_rule( + r#" + package test + p := x if { x := [1, 2, 3] } + "#, + ); + assert_no_collection_create(&program); + assert_literal_exists(&program, &Value::from_json_str("[1, 2, 3]").unwrap()); +} + +#[test] +fn constant_set_is_hoisted() { + let program = compile_rule( + r#" + package test + p := x if { x := {1, 2, 3} } + "#, + ); + assert_no_collection_create(&program); + let expected_set = Value::Set(Rc::new( + [1, 2, 3] + .into_iter() + .map(Value::from) + .collect::>(), + )); + assert_literal_exists(&program, &expected_set); +} + +#[test] +fn constant_object_is_hoisted() { + let program = compile_rule( + r#" + package test + p := x if { x := {"a": 1, "b": 2} } + "#, + ); + assert_no_collection_create(&program); + assert_literal_exists( + &program, + &Value::from_json_str(r#"{"a": 1, "b": 2}"#).unwrap(), + ); +} + +#[test] +fn nested_constant_collection_is_hoisted() { + let program = compile_rule( + r#" + package test + p := x if { x := [1, [2, 3], {"k": "v"}] } + "#, + ); + assert_no_collection_create(&program); + assert_literal_exists( + &program, + &Value::from_json_str(r#"[1, [2, 3], {"k": "v"}]"#).unwrap(), + ); +} + +#[test] +fn non_constant_array_is_not_hoisted() { + let program = compile_rule( + r#" + package test + p := x if { y := 1; x := [y, 2, 3] } + "#, + ); + // This array contains a variable reference, so it must NOT be hoisted. + let has_array_create = program + .instructions + .iter() + .any(|i| matches!(i, Instruction::ArrayCreate { .. })); + assert!( + has_array_create, + "non-constant array should use ArrayCreate" + ); +} + +// --- AssertEq fusion tests --- + +/// Count occurrences of a specific instruction pattern in the program. +fn count_instructions( + program: ®orus::rvm::program::Program, + pred: impl Fn(&Instruction) -> bool, +) -> usize { + program.instructions.iter().filter(|i| pred(i)).count() +} + +#[test] +fn equality_check_emits_assert_eq() { + // Assignment `x = 1` followed by `x = 1` triggers EqualityCheck in destructuring. + let program = compile_rule( + r#" + package test + p if { x = 1; x = 1 } + "#, + ); + let assert_eq_count = + count_instructions(&program, |i| matches!(i, Instruction::AssertEq { .. })); + assert!( + assert_eq_count > 0, + "expected AssertEq instruction for equality check" + ); +} + +#[test] +fn destructuring_equality_emits_assert_eq() { + let program = compile_rule( + r#" + package test + p if { [1, x] := [1, 2] } + "#, + ); + let assert_eq_count = + count_instructions(&program, |i| matches!(i, Instruction::AssertEq { .. })); + assert!( + assert_eq_count > 0, + "expected AssertEq for destructuring equality" + ); +} + +#[test] +fn not_expr_emits_assert_not() { + let program = compile_rule( + r#" + package test + p if { not false } + "#, + ); + let assert_not_count = + count_instructions(&program, |i| matches!(i, Instruction::AssertNot { .. })); + assert!( + assert_not_count > 0, + "expected AssertNot for `not` expression" + ); + // The Not+AssertCondition pair should be fused — no separate Not instruction. + let not_count = count_instructions(&program, |i| matches!(i, Instruction::Not { .. })); + assert_eq!(not_count, 0, "Not should be fused into AssertNot"); +} + +// --- B-11: early_exit_on_first_success flag tests --- + +/// Find a RuleInfo by name suffix (e.g., "check" matches "data.test.check"). +fn find_rule_info<'a>( + program: &'a regorus::rvm::program::Program, + name_suffix: &str, +) -> &'a regorus::rvm::program::RuleInfo { + program + .rule_infos + .iter() + .find(|ri| ri.name.ends_with(name_suffix)) + .unwrap_or_else(|| panic!("no RuleInfo ending with '{name_suffix}'")) +} + +#[test] +fn early_exit_set_for_implicit_true_multi_def() { + let program = compile_rule( + r#" + package test + p if { 1 == 1 } + p if { 2 == 2 } + "#, + ); + let ri = find_rule_info(&program, ".p"); + assert!( + ri.early_exit_on_first_success, + "two implicit-true defs should set early_exit_on_first_success" + ); +} + +#[test] +fn early_exit_set_for_same_literal_string() { + let program = compile_rule( + r#" + package test + p := "ok" if { 1 == 1 } + p := "ok" if { 2 == 2 } + "#, + ); + let ri = find_rule_info(&program, ".p"); + assert!( + ri.early_exit_on_first_success, + "two defs both returning \"ok\" should set early_exit_on_first_success" + ); +} + +#[test] +fn early_exit_not_set_for_different_literals() { + let program = compile_rule( + r#" + package test + p := "a" if { 1 == 1 } + p := "b" if { 2 == 2 } + "#, + ); + let ri = find_rule_info(&program, ".p"); + assert!( + !ri.early_exit_on_first_success, + "defs returning different literals must NOT set early_exit_on_first_success" + ); +} + +#[test] +fn early_exit_not_set_for_computed_values() { + let program = compile_rule( + r#" + package test + p := x if { x := 1 + 1 } + p := x if { x := 2 + 0 } + "#, + ); + let ri = find_rule_info(&program, ".p"); + assert!( + !ri.early_exit_on_first_success, + "computed expressions must NOT set early_exit_on_first_success" + ); +} + +#[test] +fn early_exit_not_set_for_single_definition() { + let program = compile_rule( + r#" + package test + p if { 1 == 1 } + "#, + ); + let ri = find_rule_info(&program, ".p"); + assert!( + !ri.early_exit_on_first_success, + "single definition should not set early_exit (only ≥2 defs)" + ); +} + +#[test] +fn early_exit_not_set_for_else_with_different_values() { + let program = compile_rule( + r#" + package test + p := "a" if { false } else := "b" if { true } + p := "a" if { true } + "#, + ); + let ri = find_rule_info(&program, ".p"); + assert!( + !ri.early_exit_on_first_success, + "else branches with different values must NOT set early_exit_on_first_success" + ); +} + +#[test] +fn early_exit_set_for_else_with_same_values() { + let program = compile_rule( + r#" + package test + p := "x" if { false } else := "x" if { true } + p := "x" if { true } + "#, + ); + let ri = find_rule_info(&program, ".p"); + assert!( + ri.early_exit_on_first_success, + "else branches all returning same literal should set early_exit_on_first_success" + ); +} + +#[test] +fn early_exit_set_for_implicit_true_function() { + let mut engine = Engine::new(); + engine + .add_policy( + "test.rego".to_string(), + r#" + package test + check(x) if { x > 0 } + check(x) if { x < -10 } + p := check(5) + "# + .to_string(), + ) + .expect("failed to add policy"); + let compiled = engine + .compile_with_entrypoint(&Rc::from("data.test.p")) + .expect("failed to compile"); + let program = Compiler::compile_from_policy(&compiled, &["data.test.p"]) + .expect("failed to compile to RVM"); + let ri = find_rule_info(&program, ".check"); + assert!( + ri.early_exit_on_first_success, + "implicit-true function with 2 defs should set early_exit_on_first_success" + ); +} diff --git a/tests/rvm/mod.rs b/tests/rvm/mod.rs index d28bff6..0fb26aa 100644 --- a/tests/rvm/mod.rs +++ b/tests/rvm/mod.rs @@ -1,3 +1,4 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. +mod compiler; mod rego; diff --git a/tests/rvm/rego/cases/early_exit_same_value.yaml b/tests/rvm/rego/cases/early_exit_same_value.yaml new file mode 100644 index 0000000..b18115f --- /dev/null +++ b/tests/rvm/rego/cases/early_exit_same_value.yaml @@ -0,0 +1,283 @@ +# B-11: Early exit for same-value multi-definition rules +# +# When all definitions of a Complete or function rule produce the same +# static value, the VM can stop after the first successful definition. +# These tests verify correctness: both that the optimization produces +# the right result and that edge cases (different values, else branches, +# computed expressions) remain correct. + +cases: + # ── Implicit-true multi-def function rules ──────────────────────── + + - note: implicit_true_two_definitions_first_succeeds + description: Two implicit-true definitions; first succeeds → true + modules: + - | + package test + check(x) if { x > 0 } + check(x) if { x < -10 } + p := check(5) + query: data.test.p + want_result: true + + - note: implicit_true_two_definitions_second_succeeds + description: Two implicit-true definitions; only second succeeds → true + modules: + - | + package test + check(x) if { x > 100 } + check(x) if { x < 0 } + p := check(-5) + query: data.test.p + want_result: true + + - note: implicit_true_two_definitions_neither_succeeds + description: Two implicit-true definitions; neither succeeds → undefined + modules: + - | + package test + check(x) if { x > 100 } + check(x) if { x < -100 } + p := check(5) + query: data.test.p + want_result: "#undefined" + + - note: implicit_true_four_definitions + description: Four implicit-true defs (like mountSource_ok); third succeeds + modules: + - | + package test + validate(x) if { x == "a" } + validate(x) if { x == "b" } + validate(x) if { x == "c" } + validate(x) if { x == "d" } + p := validate("c") + query: data.test.p + want_result: true + + - note: implicit_true_complete_rule_two_defs + description: Complete rule with two implicit-true defs + modules: + - | + package test + allowed if { input.role == "admin" } + allowed if { input.role == "superuser" } + p := allowed + input: {"role": "superuser"} + query: data.test.p + want_result: true + + # ── Same literal value (non-true) across definitions ─────────────── + + - note: same_string_value_two_defs + description: Two defs returning same string constant + modules: + - | + package test + label(x) := "ok" if { x > 0 } + label(x) := "ok" if { x == 0 } + p := label(0) + query: data.test.p + want_result: "ok" + + - note: same_number_value_two_defs + description: Two defs returning same numeric constant + modules: + - | + package test + code(x) := 42 if { x == "answer" } + code(x) := 42 if { x == "the answer" } + p := code("the answer") + query: data.test.p + want_result: 42 + + - note: same_bool_false_two_defs + description: Two defs returning explicit false + modules: + - | + package test + deny(x) := false if { x == "blocked" } + deny(x) := false if { x == "banned" } + p := deny("banned") + query: data.test.p + want_result: false + + # ── Different values across definitions (NO early exit) ──────────── + + - note: different_string_values_first_succeeds + description: Two defs with different strings; first succeeds → its value + modules: + - | + package test + classify(x) := "positive" if { x > 0 } + classify(x) := "non-positive" if { x <= 0 } + p := classify(5) + query: data.test.p + want_result: "positive" + + - note: different_string_values_second_succeeds + description: Two defs with different strings; second succeeds → its value + modules: + - | + package test + classify(x) := "positive" if { x > 0 } + classify(x) := "non-positive" if { x <= 0 } + p := classify(-3) + query: data.test.p + want_result: "non-positive" + + # ── Else branches ────────────────────────────────────────────────── + + - note: else_same_value_across_defs + description: Two defs, each with else, all branches return same value + modules: + - | + package test + result := "match" if { + input.x > 10 + } + else := "match" if { + input.x > 5 + } + result := "match" if { + input.y > 10 + } + else := "match" if { + input.y > 5 + } + p := result + input: {"x": 1, "y": 7} + query: data.test.p + want_result: "match" + + - note: else_different_values_within_def + description: One def with else returning different value → no early exit + modules: + - | + package test + grade(x) := "A" if { + x >= 90 + } + else := "B" if { + x >= 80 + } + grade(x) := "C" if { x >= 70; x < 80 } + p := grade(85) + query: data.test.p + want_result: "B" + + - note: else_different_values_within_def_second + description: Else chain where second definition fires + modules: + - | + package test + grade(x) := "A" if { + x >= 90 + } + else := "B" if { + x >= 80 + } + grade(x) := "C" if { x >= 70; x < 80 } + p := grade(75) + query: data.test.p + want_result: "C" + + # ── Computed (non-literal) values → no early exit ────────────────── + + - note: computed_value_two_defs + description: Defs with computed expressions → no early exit, still correct + modules: + - | + package test + double(x) := x * 2 if { x > 0 } + double(x) := x * 2 if { x < 0 } + p := double(-3) + query: data.test.p + want_result: -6 + + # ── Single definition (flag doesn't matter) ──────────────────────── + + - note: single_definition_implicit_true + description: Single implicit-true def — flag not set but works fine + modules: + - | + package test + ok(x) if { x > 0 } + p := ok(5) + query: data.test.p + want_result: true + + # ── Mixed implicit-true and explicit-true ─────────────────────────── + + - note: mixed_implicit_and_explicit_true + description: One def has implicit true, another has explicit := true + modules: + - | + package test + valid(x) if { x > 0 } + valid(x) := true if { x == 0 } + p := valid(0) + query: data.test.p + want_result: true + + # ── Default value interaction ────────────────────────────────────── + + - note: implicit_true_with_default + description: Multi-def rule with default; no def succeeds → default + modules: + - | + package test + default allowed := false + allowed if { input.role == "admin" } + allowed if { input.role == "superuser" } + p := allowed + input: {"role": "viewer"} + query: data.test.p + want_result: false + + - note: implicit_true_with_default_succeeds + description: Multi-def rule with default; one def succeeds → true + modules: + - | + package test + default allowed := false + allowed if { input.role == "admin" } + allowed if { input.role == "superuser" } + p := allowed + input: {"role": "admin"} + query: data.test.p + want_result: true + + # ── Nested function calls with early exit ────────────────────────── + + - note: nested_early_exit_functions + description: Outer function calls inner multi-def function + modules: + - | + package test + inner_ok(x) if { x == "a" } + inner_ok(x) if { x == "b" } + inner_ok(x) if { x == "c" } + outer(x, y) if { + inner_ok(x) + inner_ok(y) + } + p := outer("a", "c") + query: data.test.p + want_result: true + + - note: nested_early_exit_functions_fail + description: Outer function calls inner multi-def function, inner fails + modules: + - | + package test + inner_ok(x) if { x == "a" } + inner_ok(x) if { x == "b" } + inner_ok(x) if { x == "c" } + outer(x, y) if { + inner_ok(x) + inner_ok(y) + } + p := outer("a", "d") + query: data.test.p + want_result: "#undefined" diff --git a/tests/rvm/rego/cases/objects.yaml b/tests/rvm/rego/cases/objects.yaml index 343c179..470361e 100644 --- a/tests/rvm/rego/cases/objects.yaml +++ b/tests/rvm/rego/cases/objects.yaml @@ -106,3 +106,21 @@ cases: } query: data.test.main want_result: {"username": "alice123", "user_age": 25} + + - note: object_membership_checks_values_not_keys + data: {} + modules: + - | + package test + main := "foo" in {"foo": "bar"} + query: data.test.main + want_result: false + + - note: object_membership_finds_value + data: {} + modules: + - | + package test + main := "bar" in {"foo": "bar"} + query: data.test.main + want_result: true diff --git a/tests/rvm/vm/suites/type_errors.yaml b/tests/rvm/vm/suites/type_errors.yaml index 1181391..d6796d5 100644 --- a/tests/rvm/vm/suites/type_errors.yaml +++ b/tests/rvm/vm/suites/type_errors.yaml @@ -96,15 +96,15 @@ cases: want_error: "#undefined" - note: logical_not_int - description: NOT with int operand should error - example_rego: "!42" + description: NOT with non-boolean defined operand should yield false + example_rego: "not 42" literals: - 42 instructions: - "Load { dest: 0, literal_idx: 0 }" - "Not { dest: 1, operand: 0 }" - "Return { value: 1 }" - want_error: "#undefined" + want_result: false # Invalid indexing operations - note: index_int_with_string_key