use crate::Result;
use crate::frame::Frame;
use crate::module_system::AccessCheckResult;
use ristretto_classloader::Class;
use std::sync::Arc;
#[inline]
pub(crate) fn check_class_access(frame: &Frame, target_class: &Arc<Class>) -> Result<()> {
let source_class = frame.class();
if Arc::ptr_eq(source_class, target_class) {
return Ok(());
}
let source_module = source_class.module_name().ok().flatten();
let target_module = target_class.module_name().ok().flatten();
if source_module == target_module {
return Ok(());
}
let thread = frame.thread()?;
let vm = thread.vm()?;
let result = vm.module_system().check_access(
source_module.as_deref(),
target_module.as_deref(),
target_class.name(),
);
if result.is_allowed() {
return Ok(());
}
if should_enforce_access(source_module.as_deref(), target_module.as_deref()) {
vm.module_system().require_access(
source_module.as_deref(),
target_module.as_deref(),
target_class.name(),
)?;
}
Ok(())
}
#[inline]
pub(crate) fn check_reflection_access(frame: &Frame, target_class: &Arc<Class>) -> Result<()> {
let source_class = frame.class();
if Arc::ptr_eq(source_class, target_class) {
return Ok(());
}
let source_module = source_class.module_name().ok().flatten();
let target_module = target_class.module_name().ok().flatten();
if source_module == target_module {
return Ok(());
}
let thread = frame.thread()?;
let vm = thread.vm()?;
let result = vm.module_system().check_reflection_access(
source_module.as_deref(),
target_module.as_deref(),
target_class.name(),
);
if result.is_allowed() {
return Ok(());
}
if should_enforce_access(source_module.as_deref(), target_module.as_deref()) {
vm.module_system().require_reflection_access(
source_module.as_deref(),
target_module.as_deref(),
target_class.name(),
)?;
}
Ok(())
}
#[inline]
pub(crate) fn check_class_access_by_name(
frame: &Frame,
target_class_name: &str,
target_module: Option<&str>,
) -> Result<()> {
let source_class = frame.class();
let source_module = source_class.module_name().ok().flatten();
if source_module.as_deref() == target_module {
return Ok(());
}
let thread = frame.thread()?;
let vm = thread.vm()?;
let result =
vm.module_system()
.check_access(source_module.as_deref(), target_module, target_class_name);
if result.is_allowed() {
return Ok(());
}
if should_enforce_access(source_module.as_deref(), target_module) {
vm.module_system().require_access(
source_module.as_deref(),
target_module,
target_class_name,
)?;
}
Ok(())
}
#[inline]
pub(crate) fn check_reflection_access_by_name(
frame: &Frame,
target_class_name: &str,
target_module: Option<&str>,
) -> Result<()> {
let source_class = frame.class();
let source_module = source_class.module_name().ok().flatten();
if source_module.as_deref() == target_module {
return Ok(());
}
let thread = frame.thread()?;
let vm = thread.vm()?;
let result = vm.module_system().check_reflection_access(
source_module.as_deref(),
target_module,
target_class_name,
);
if result.is_allowed() {
return Ok(());
}
if should_enforce_access(source_module.as_deref(), target_module) {
vm.module_system().require_reflection_access(
source_module.as_deref(),
target_module,
target_class_name,
)?;
}
Ok(())
}
#[must_use]
pub fn access_denied_error(
result: AccessCheckResult,
source_module: Option<&str>,
target_module: Option<&str>,
target_class: &str,
) -> crate::Error {
use crate::JavaError::{IllegalAccessError, InaccessibleObjectException};
use crate::module_system::{ALL_UNNAMED, ModuleSystem};
let from = source_module.unwrap_or(ALL_UNNAMED);
let to = target_module.unwrap_or(ALL_UNNAMED);
let error_msg = ModuleSystem::illegal_access_error(from, to, target_class, result);
match result {
AccessCheckResult::NotOpened => InaccessibleObjectException(error_msg).into(),
_ => IllegalAccessError(error_msg).into(),
}
}
#[inline]
fn should_enforce_access(source_module: Option<&str>, target_module: Option<&str>) -> bool {
if source_module.is_none() || target_module.is_none() {
return false;
}
let target = target_module.unwrap_or("");
if target.starts_with("java.")
|| target.starts_with("jdk.")
|| target.starts_with("sun.")
|| target.starts_with("com.sun.")
{
return false;
}
true
}
#[cfg(test)]
mod tests {
use super::*;
use crate::module_system::ALL_UNNAMED;
#[tokio::test]
async fn test_check_class_access_same_class() -> Result<()> {
let (_vm, _thread, frame) = crate::test::frame().await?;
let class = frame.class().clone();
check_class_access(&frame, &class)?;
Ok(())
}
#[tokio::test]
async fn test_check_reflection_access_same_class() -> Result<()> {
let (_vm, _thread, frame) = crate::test::frame().await?;
let class = frame.class().clone();
check_reflection_access(&frame, &class)?;
Ok(())
}
#[tokio::test]
async fn test_check_class_access_java_lang_object() -> Result<()> {
let (_vm, thread, frame) = crate::test::frame().await?;
let object_class = thread.class("java/lang/Object").await?;
check_class_access(&frame, &object_class)?;
Ok(())
}
#[tokio::test]
async fn test_check_reflection_access_java_lang_object() -> Result<()> {
let (_vm, thread, frame) = crate::test::frame().await?;
let object_class = thread.class("java/lang/Object").await?;
check_reflection_access(&frame, &object_class)?;
Ok(())
}
#[tokio::test]
async fn test_check_class_access_by_name_same_module() -> Result<()> {
let (_vm, _thread, frame) = crate::test::frame().await?;
check_class_access_by_name(&frame, "com/example/MyClass", None)?;
Ok(())
}
#[tokio::test]
async fn test_check_reflection_access_by_name_same_module() -> Result<()> {
let (_vm, _thread, frame) = crate::test::frame().await?;
check_reflection_access_by_name(&frame, "com/example/MyClass", None)?;
Ok(())
}
#[tokio::test]
async fn test_check_class_access_to_system_module() -> Result<()> {
let (_vm, thread, frame) = crate::test::frame().await?;
let string_class = thread.class("java/lang/String").await?;
check_class_access(&frame, &string_class)?;
Ok(())
}
#[tokio::test]
async fn test_check_class_access_multiple_system_classes() -> Result<()> {
let (_vm, thread, frame) = crate::test::frame().await?;
let classes = ["java/lang/Integer", "java/util/ArrayList", "java/io/File"];
for class_name in &classes {
let class = thread.class(class_name).await?;
check_class_access(&frame, &class)?;
}
Ok(())
}
#[test]
fn test_should_enforce_access_unnamed_modules() {
assert!(!should_enforce_access(None, None));
assert!(!should_enforce_access(None, Some("app.module")));
assert!(!should_enforce_access(Some("app.module"), None));
}
#[test]
fn test_should_enforce_access_system_modules() {
let source = Some("my.app");
assert!(!should_enforce_access(source, Some("java.base")));
assert!(!should_enforce_access(source, Some("java.sql")));
assert!(!should_enforce_access(source, Some("jdk.compiler")));
assert!(!should_enforce_access(source, Some("sun.misc")));
assert!(!should_enforce_access(
source,
Some("com.sun.crypto.provider")
));
}
#[test]
fn test_should_enforce_access_application_modules() {
let source = Some("my.app");
let target = Some("other.app");
assert!(should_enforce_access(source, target));
}
#[test]
fn test_access_denied_error_not_readable() {
let error = access_denied_error(
AccessCheckResult::NotReadable,
Some("my.app"),
Some("other.app"),
"other/internal/Secret",
);
let error_str = format!("{error:?}");
assert!(error_str.contains("IllegalAccessError"));
}
#[test]
fn test_access_denied_error_not_exported() {
let error = access_denied_error(
AccessCheckResult::NotExported,
Some("my.app"),
Some("other.app"),
"other/internal/Secret",
);
let error_str = format!("{error:?}");
assert!(error_str.contains("IllegalAccessError"));
}
#[test]
fn test_access_denied_error_not_opened() {
let error = access_denied_error(
AccessCheckResult::NotOpened,
Some("my.app"),
Some("other.app"),
"other/internal/Secret",
);
let error_str = format!("{error:?}");
assert!(error_str.contains("InaccessibleObjectException"));
}
#[test]
fn test_access_denied_error_unnamed_module() {
let error = access_denied_error(
AccessCheckResult::NotExported,
None,
Some("java.base"),
"java/lang/internal/Secret",
);
let error_str = format!("{error:?}");
assert!(error_str.contains("unnamed module") || error_str.contains(ALL_UNNAMED));
}
}