use crate::Result;
use crate::frame::Frame;
use crate::method_ref_cache::{
InvokeKind, MethodRefError, MethodRefErrorKind, MethodRefKey, ResolvedMethodRef,
};
use crate::module_system::{ALL_UNNAMED, AccessCheckResult, ModuleSystem};
use ristretto_classfile::Constant;
use ristretto_classloader::{Class, Method};
use std::sync::Arc;
#[derive(Clone)]
pub struct MethodResolution {
pub declaring_class: Arc<Class>,
pub method: Arc<Method>,
pub method_name: String,
pub method_descriptor: String,
pub is_polymorphic: bool,
pub param_count: usize,
pub has_return_type: bool,
}
pub async fn resolve_method_ref(
frame: &Frame,
method_index: u16,
invoke_kind: InvokeKind,
) -> Result<MethodResolution> {
let thread = frame.thread()?;
let vm = thread.vm()?;
let caller_class = frame.class();
let cache_key = MethodRefKey::new(caller_class.name().to_string(), method_index);
if let Some(result) = vm.method_ref_cache().get(&cache_key) {
let resolved = result?;
return Ok(MethodResolution {
declaring_class: resolved.declaring_class.clone(),
method: resolved.method.clone(),
method_name: resolved.method_name.clone(),
method_descriptor: resolved.method_descriptor.clone(),
is_polymorphic: resolved.is_polymorphic,
param_count: resolved.param_count,
has_return_type: resolved.has_return_type,
});
}
let constant_pool = caller_class.constant_pool();
let constant = constant_pool.try_get(method_index)?;
let (class_index, name_and_type_index, is_interface_method) = match constant {
Constant::MethodRef {
class_index,
name_and_type_index,
} => (*class_index, *name_and_type_index, false),
Constant::InterfaceMethodRef {
class_index,
name_and_type_index,
} => (*class_index, *name_and_type_index, true),
_ => {
return Err(
ristretto_classfile::Error::InvalidConstantPoolIndexType(method_index).into(),
);
}
};
let class_name = constant_pool.try_get_class(class_index)?;
let target_class = thread.class_java_str(class_name).await?;
let class_name = class_name.to_str_lossy();
validate_class_kind(&target_class, invoke_kind, is_interface_method, &class_name)?;
check_jpms_access(frame, &target_class)?;
let (name_index, descriptor_index) =
constant_pool.try_get_name_and_type(name_and_type_index)?;
let method_name = constant_pool.try_get_utf8(*name_index)?;
let method_name = method_name.to_str_lossy();
let method_descriptor = constant_pool.try_get_utf8(*descriptor_index)?;
let method_descriptor = method_descriptor.to_str_lossy();
let (resolved_class, method) = if invoke_kind == InvokeKind::Interface {
lookup_interface_method(&target_class, &method_name, &method_descriptor)?
} else {
match lookup_method(&target_class, &method_name, &method_descriptor) {
Ok(result) => result,
Err(e) => {
if is_holder_class_for_resolution(&class_name) {
let vm = thread.vm()?;
let registry = vm.method_registry();
if registry
.method(&class_name, &method_name, &method_descriptor)
.is_some()
{
let synthetic_method = create_synthetic_intrinsic_method(
&class_name,
&method_name,
&method_descriptor,
)?;
(target_class.clone(), synthetic_method)
} else {
return Err(e);
}
} else {
return Err(e);
}
}
}
};
validate_method_for_invoke(&method, &method_name, &method_descriptor, invoke_kind)?;
let resolved_ref = ResolvedMethodRef::new(
resolved_class.clone(),
method.clone(),
invoke_kind,
method_descriptor.to_string(),
);
let is_polymorphic = resolved_ref.is_polymorphic;
let param_count = resolved_ref.param_count;
let has_return_type = resolved_ref.has_return_type;
vm.method_ref_cache()
.store_resolved(cache_key, resolved_ref);
Ok(MethodResolution {
declaring_class: resolved_class,
method,
method_name: method_name.to_string(),
method_descriptor: method_descriptor.to_string(),
is_polymorphic,
param_count,
has_return_type,
})
}
fn validate_class_kind(
target_class: &Arc<Class>,
invoke_kind: InvokeKind,
is_interface_method: bool,
class_name: &str,
) -> Result<()> {
use crate::JavaError::IncompatibleClassChangeError;
match invoke_kind {
InvokeKind::Static => {
if is_interface_method && !target_class.is_interface() {
return Err(IncompatibleClassChangeError(format!(
"Expected interface, found class: {class_name}"
))
.into());
}
if !is_interface_method && target_class.is_interface() {
return Err(IncompatibleClassChangeError(format!(
"Expected class, found interface: {class_name}"
))
.into());
}
}
InvokeKind::Interface => {
if !target_class.is_interface() {
return Err(IncompatibleClassChangeError(format!(
"{class_name} is not an interface"
))
.into());
}
}
InvokeKind::Virtual | InvokeKind::Special => {
}
}
Ok(())
}
fn validate_method_for_invoke(
method: &Method,
method_name: &str,
method_descriptor: &str,
invoke_kind: InvokeKind,
) -> Result<()> {
use crate::JavaError::IncompatibleClassChangeError;
match invoke_kind {
InvokeKind::Static => {
if !method.is_static() {
return Err(IncompatibleClassChangeError(format!(
"Method {method_name}{method_descriptor} is not static"
))
.into());
}
}
InvokeKind::Virtual | InvokeKind::Interface => {
if method.is_static() {
return Err(IncompatibleClassChangeError(format!(
"Method {method_name}{method_descriptor} is static"
))
.into());
}
}
InvokeKind::Special => {
}
}
Ok(())
}
pub fn check_jpms_access(frame: &Frame, target_class: &Arc<Class>) -> Result<()> {
let caller_class = frame.class();
if Arc::ptr_eq(caller_class, target_class) {
return Ok(());
}
let caller_module = caller_class.module_name().ok().flatten();
let target_module = target_class.module_name().ok().flatten();
if caller_module == target_module {
return Ok(());
}
let thread = frame.thread()?;
let vm = thread.vm()?;
let result = vm.module_system().check_access(
caller_module.as_deref(),
target_module.as_deref(),
target_class.name(),
);
if result.is_allowed() {
return Ok(());
}
if !should_enforce_jpms_access(caller_module.as_deref(), target_module.as_deref()) {
return Ok(());
}
let from = caller_module.as_deref().unwrap_or(ALL_UNNAMED);
let to = target_module.as_deref().unwrap_or(ALL_UNNAMED);
let error_msg = ModuleSystem::illegal_access_error(from, to, target_class.name(), result);
Err(crate::JavaError::IllegalAccessError(error_msg).into())
}
fn should_enforce_jpms_access(caller_module: Option<&str>, target_module: Option<&str>) -> bool {
if caller_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
}
pub fn lookup_method(
class: &Arc<Class>,
name: &str,
descriptor: &str,
) -> Result<(Arc<Class>, Arc<Method>)> {
if let Some(method) = class.method(name, descriptor) {
return Ok((class.clone(), method));
}
let mut current = class.parent()?;
while let Some(parent) = current {
if let Some(method) = parent.method(name, descriptor) {
return Ok((parent, method));
}
current = parent.parent()?;
}
let mut interfaces_to_check: Vec<Arc<Class>> = class.interfaces()?;
let mut visited = std::collections::HashSet::new();
visited.insert(class.name().to_string());
let mut abstract_method: Option<(Arc<Class>, Arc<Method>)> = None;
while let Some(interface) = interfaces_to_check.pop() {
if !visited.insert(interface.name().to_string()) {
continue;
}
if let Some(method) = interface.method(name, descriptor) {
if !method.is_abstract() {
return Ok((interface, method));
} else if abstract_method.is_none() {
abstract_method = Some((interface.clone(), method));
}
}
interfaces_to_check.extend(interface.interfaces()?);
}
let mut class_to_check = class.parent()?;
while let Some(parent_class) = class_to_check {
let mut parent_interfaces: Vec<Arc<Class>> = parent_class.interfaces()?;
while let Some(interface) = parent_interfaces.pop() {
if !visited.insert(interface.name().to_string()) {
continue;
}
if let Some(method) = interface.method(name, descriptor) {
if !method.is_abstract() {
return Ok((interface, method));
} else if abstract_method.is_none() {
abstract_method = Some((interface.clone(), method));
}
}
parent_interfaces.extend(interface.interfaces()?);
}
class_to_check = parent_class.parent()?;
}
if let Some((interface, method)) = abstract_method {
return Ok((interface, method));
}
Err(crate::JavaError::NoSuchMethodError(format!(
"Method {name}{descriptor} not found in class {}",
class.name()
))
.into())
}
pub fn lookup_interface_method(
interface: &Arc<Class>,
name: &str,
descriptor: &str,
) -> Result<(Arc<Class>, Arc<Method>)> {
if let Some(method) = interface.method(name, descriptor) {
return Ok((interface.clone(), method));
}
let mut interfaces_to_check: Vec<Arc<Class>> = interface.interfaces()?;
let mut visited = std::collections::HashSet::new();
visited.insert(interface.name().to_string());
while let Some(super_interface) = interfaces_to_check.pop() {
if !visited.insert(super_interface.name().to_string()) {
continue;
}
if let Some(method) = super_interface.method(name, descriptor) {
return Ok((super_interface, method));
}
interfaces_to_check.extend(super_interface.interfaces()?);
}
if let Ok(Some(object_class)) = interface.parent()
&& let Some(method) = object_class.method(name, descriptor)
{
return Ok((object_class, method));
}
Err(crate::JavaError::NoSuchMethodError(format!(
"Method {name}{descriptor} not found in interface {}",
interface.name()
))
.into())
}
#[must_use]
pub fn create_jpms_error(
result: AccessCheckResult,
caller_module: Option<&str>,
target_module: Option<&str>,
target_class: &str,
) -> MethodRefError {
let from = caller_module.unwrap_or(ALL_UNNAMED);
let to = target_module.unwrap_or(ALL_UNNAMED);
let message = ModuleSystem::illegal_access_error(from, to, target_class, result);
let kind = match result {
AccessCheckResult::NotReadable => MethodRefErrorKind::ModuleNotReadable,
AccessCheckResult::NotExported | AccessCheckResult::NotOpened => {
MethodRefErrorKind::PackageNotExported
}
AccessCheckResult::Allowed => MethodRefErrorKind::InternalError, };
MethodRefError::new(kind, message)
}
fn is_holder_class_for_resolution(class_name: &str) -> bool {
let normalized = class_name.replace('.', "/");
matches!(
normalized.as_str(),
"java/lang/invoke/DirectMethodHandle$Holder"
| "java/lang/invoke/DelegatingMethodHandle$Holder"
| "java/lang/invoke/Invokers$Holder"
| "java/lang/invoke/LambdaForm$Holder"
| "java/lang/invoke/VarHandleGuards"
) || normalized.starts_with("java/lang/invoke/LambdaForm$")
}
fn create_synthetic_intrinsic_method(
_class_name: &str,
method_name: &str,
method_descriptor: &str,
) -> Result<Arc<Method>> {
use ristretto_classfile::MethodAccessFlags;
let definition = ristretto_classfile::Method {
access_flags: MethodAccessFlags::PUBLIC
| MethodAccessFlags::STATIC
| MethodAccessFlags::NATIVE,
name_index: 0, descriptor_index: 0, attributes: Vec::new(),
};
let method_descriptor = ristretto_classfile::JavaStr::cow_from_str(method_descriptor);
let (parameters, return_type) =
ristretto_classfile::FieldType::parse_method_descriptor(&method_descriptor)?;
let method = Method::new_synthetic(
definition,
method_name.to_string(),
method_descriptor.to_string(),
parameters,
return_type,
);
Ok(Arc::new(method))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::VM;
#[test]
fn test_should_enforce_jpms_access_unnamed() {
assert!(!should_enforce_jpms_access(None, None));
assert!(!should_enforce_jpms_access(None, Some("app.module")));
assert!(!should_enforce_jpms_access(Some("app.module"), None));
}
#[test]
fn test_should_enforce_jpms_access_system_modules() {
let caller = Some("my.app");
assert!(!should_enforce_jpms_access(caller, Some("java.base")));
assert!(!should_enforce_jpms_access(caller, Some("java.sql")));
assert!(!should_enforce_jpms_access(caller, Some("jdk.compiler")));
assert!(!should_enforce_jpms_access(caller, Some("sun.misc")));
assert!(!should_enforce_jpms_access(
caller,
Some("com.sun.crypto.provider")
));
}
#[test]
fn test_should_enforce_jpms_access_app_modules() {
assert!(should_enforce_jpms_access(
Some("my.app"),
Some("other.app")
));
assert!(should_enforce_jpms_access(
Some("com.example"),
Some("org.lib")
));
}
#[test]
fn test_create_jpms_error_not_readable() {
let error = create_jpms_error(
AccessCheckResult::NotReadable,
Some("my.app"),
Some("other.app"),
"other/api/Service",
);
assert_eq!(error.kind, MethodRefErrorKind::ModuleNotReadable);
assert!(error.message.contains("does not read"));
}
#[test]
fn test_create_jpms_error_not_exported() {
let error = create_jpms_error(
AccessCheckResult::NotExported,
Some("my.app"),
Some("other.app"),
"other/internal/Secret",
);
assert_eq!(error.kind, MethodRefErrorKind::PackageNotExported);
assert!(error.message.contains("does not export"));
}
#[tokio::test]
async fn test_lookup_method_found_in_class() -> Result<()> {
let vm = VM::default().await?;
let class = vm.class("java.lang.String").await?;
let (resolved_class, method) = lookup_method(&class, "length", "()I")?;
assert_eq!(resolved_class.name(), "java/lang/String");
assert_eq!(method.name(), "length");
Ok(())
}
#[tokio::test]
async fn test_lookup_method_found_in_superclass() -> Result<()> {
let vm = VM::default().await?;
let class = vm.class("java.util.ArrayList").await?;
let (resolved_class, method) = lookup_method(&class, "toString", "()Ljava/lang/String;")?;
assert_eq!(resolved_class.name(), "java/util/AbstractCollection");
assert_eq!(method.name(), "toString");
Ok(())
}
#[tokio::test]
async fn test_lookup_method_not_found() -> Result<()> {
let vm = VM::default().await?;
let class = vm.class("java.lang.String").await?;
let result = lookup_method(&class, "nonExistentMethod", "()V");
assert!(result.is_err());
Ok(())
}
#[tokio::test]
async fn test_validate_class_kind_static_with_class() -> Result<()> {
let vm = VM::default().await?;
let class = vm.class("java.lang.String").await?;
let result = validate_class_kind(&class, InvokeKind::Static, false, "java/lang/String");
assert!(result.is_ok());
Ok(())
}
#[tokio::test]
async fn test_validate_class_kind_interface_requires_interface() -> Result<()> {
let vm = VM::default().await?;
let class = vm.class("java.lang.String").await?;
let result = validate_class_kind(&class, InvokeKind::Interface, true, "java/lang/String");
assert!(result.is_err());
Ok(())
}
#[tokio::test]
async fn test_validate_method_for_invoke_static() -> Result<()> {
let vm = VM::default().await?;
let class = vm.class("java.lang.Integer").await?;
if let Some(method) = class.method("valueOf", "(I)Ljava/lang/Integer;") {
let result = validate_method_for_invoke(
&method,
"valueOf",
"(I)Ljava/lang/Integer;",
InvokeKind::Static,
);
assert!(result.is_ok());
}
if let Some(method) = class.method("intValue", "()I") {
let result = validate_method_for_invoke(&method, "intValue", "()I", InvokeKind::Static);
assert!(result.is_err());
}
Ok(())
}
#[tokio::test]
async fn test_validate_method_for_invoke_virtual() -> Result<()> {
let vm = VM::default().await?;
let class = vm.class("java.lang.Integer").await?;
if let Some(method) = class.method("intValue", "()I") {
let result =
validate_method_for_invoke(&method, "intValue", "()I", InvokeKind::Virtual);
assert!(result.is_ok());
}
if let Some(method) = class.method("valueOf", "(I)Ljava/lang/Integer;") {
let result = validate_method_for_invoke(
&method,
"valueOf",
"(I)Ljava/lang/Integer;",
InvokeKind::Virtual,
);
assert!(result.is_err());
}
Ok(())
}
}