use std::cell::{Ref, RefCell};
use std::rc::Rc;
use crate::arch::Arch;
use crate::encoding::dataflows::MemoryAccesses;
use crate::oracle::MappableArea;
use crate::state::random::{update_memory_addresses_unchecked, StateGen};
use crate::state::{AsSystemState, StateByte, SystemState, SystemStateByteView};
#[derive(Copy, Clone)]
struct ComplexMappingEntry {
byte: StateByte,
from: u8,
to: u8,
}
impl ComplexMappingEntry {
fn apply<A: Arch>(&self, state: &mut SystemState<A>, view: &SystemStateByteView<A>) {
view.set(state, self.byte, self.to);
}
fn revert<A: Arch>(&self, state: &mut SystemState<A>, view: &SystemStateByteView<A>) {
view.set(state, self.byte, self.from);
}
}
enum ComplexMapping {
Single(ComplexMappingEntry),
Multiple(Vec<ComplexMappingEntry>),
}
#[derive(Clone)]
pub struct ComplexJitState<'a, A: Arch> {
data: Rc<RefCell<ComplexState<'a, A>>>,
id: usize,
}
struct ComplexItem {
mapping: ComplexMapping,
need_address_update: bool,
}
impl ComplexItem {
fn apply<A: Arch>(&self, state: &mut SystemState<A>, view: &SystemStateByteView<A>) {
match &self.mapping {
ComplexMapping::Single(item) => item.apply(state, view),
ComplexMapping::Multiple(items) => {
for item in items.iter() {
item.apply(state, view)
}
},
}
}
fn revert<A: Arch>(&self, state: &mut SystemState<A>, view: &SystemStateByteView<A>) {
match &self.mapping {
ComplexMapping::Single(item) => item.revert(state, view),
ComplexMapping::Multiple(items) => {
for item in items.iter() {
item.revert(state, view)
}
},
}
}
}
struct ComplexState<'a, A: Arch> {
state: SystemState<A>,
active: usize,
items: Vec<ComplexItem>,
view: SystemStateByteView<'a, A>,
accesses: &'a MemoryAccesses<A>,
}
pub struct ComplexStateRef<'a, A: Arch>(Ref<'a, ComplexState<'a, A>>);
impl<A: Arch> AsRef<SystemState<A>> for ComplexStateRef<'_, A> {
fn as_ref(&self) -> &SystemState<A> {
self.0.state.as_ref()
}
}
impl<A: Arch> AsSystemState<A> for ComplexJitState<'_, A> {
type Output<'a>
= ComplexStateRef<'a, A>
where
Self: 'a;
fn as_system_state(&self) -> Self::Output<'_> {
if self.data.borrow().active != self.id {
let data = &mut *self.data.borrow_mut();
let state = &mut data.state;
if data.active != usize::MAX {
data.items[data.active].revert(state, &data.view);
}
data.items[self.id].apply(state, &data.view);
if (data.active != usize::MAX && data.items[data.active].need_address_update)
|| data.items[self.id].need_address_update
{
update_memory_addresses_unchecked(data.accesses, state);
}
data.active = self.id;
}
ComplexStateRef(self.data.borrow())
}
fn num_memory_mappings(&self) -> usize {
self.data.borrow().state.memory().len()
}
}
pub struct ComplexJitStateBuilder<'a, 's, A: Arch, M: MappableArea> {
data: Rc<RefCell<ComplexState<'a, A>>>,
state_gen: &'s StateGen<'a, A, M>,
view: SystemStateByteView<'a, A>,
}
impl<'a, A: Arch> ComplexJitState<'a, A> {
pub fn build<'s, M: MappableArea>(
base_state: SystemState<A>, state_gen: &'s StateGen<'a, A, M>, view: SystemStateByteView<'a, A>,
) -> ComplexJitStateBuilder<'a, 's, A, M> {
ComplexJitStateBuilder {
data: Rc::new(RefCell::new(ComplexState {
state: base_state,
active: usize::MAX,
items: Vec::new(),
view,
accesses: state_gen.accesses,
})),
state_gen,
view,
}
}
}
impl<'a, A: Arch, M: MappableArea> ComplexJitStateBuilder<'a, '_, A, M> {
pub fn as_original_system_state(&self) -> ComplexStateRef<A> {
{
let data = &mut *self.data.borrow_mut();
let s = &mut data.state;
if data.active != usize::MAX {
data.items[data.active].revert(s, &self.view);
if data.items[data.active].need_address_update {
update_memory_addresses_unchecked(data.accesses, s);
}
data.active = usize::MAX;
}
}
ComplexStateRef(self.data.borrow())
}
pub fn create(
&self, bytes_affected: &[StateByte], create: impl FnOnce(&mut SystemState<A>) -> bool,
) -> Option<ComplexJitState<'a, A>> {
let need_address_update = self.state_gen.needs_adapt_from_bytes(&self.view, bytes_affected);
let data = &mut *self.data.borrow_mut();
let id = data.items.len();
let s = &mut data.state;
if data.active != usize::MAX {
data.items[data.active].revert(s, &self.view);
if data.items[data.active].need_address_update {
update_memory_addresses_unchecked(data.accesses, s);
}
data.active = usize::MAX;
}
let mut item = ComplexItem {
mapping: if bytes_affected.len() == 1 {
let byte = bytes_affected[0];
ComplexMapping::Single(ComplexMappingEntry {
byte,
from: self.view.get(s, byte),
to: 0,
})
} else {
ComplexMapping::Multiple(
bytes_affected
.iter()
.map(|&byte| ComplexMappingEntry {
byte,
from: self.view.get(s, byte),
to: 0,
})
.collect(),
)
},
need_address_update,
};
if !create(s) {
item.revert(s, &self.view);
if need_address_update {
update_memory_addresses_unchecked(data.accesses, s);
}
None
} else if !need_address_update || self.state_gen.adapt(s, false) {
match &mut item.mapping {
ComplexMapping::Single(item) => item.to = self.view.get(s, item.byte),
ComplexMapping::Multiple(items) => {
for item in items.iter_mut() {
item.to = self.view.get(s, item.byte)
}
items.retain(|item| item.from != item.to);
},
}
data.active = id;
data.items.push(item);
Some(ComplexJitState {
data: self.data.clone(),
id,
})
} else {
item.revert(s, &self.view);
if need_address_update {
update_memory_addresses_unchecked(data.accesses, s);
}
None
}
}
pub fn create_with_change_list(
&self, bytes_affected: &[StateByte], mut new_values: impl Iterator<Item = u8>,
) -> Option<ComplexJitState<'a, A>> {
if self.state_gen.needs_adapt_from_bytes(&self.view, bytes_affected) {
self.create(bytes_affected, move |state| {
for (&b, v) in bytes_affected.iter().zip(new_values) {
self.view.set(state, b, v);
}
true
})
} else {
let data = &mut *self.data.borrow_mut();
let id = data.items.len();
let s = &mut data.state;
if data.active != usize::MAX {
data.items[data.active].revert(s, &self.view);
if data.items[data.active].need_address_update {
update_memory_addresses_unchecked(data.accesses, s);
}
data.active = usize::MAX;
}
let item = ComplexItem {
mapping: if bytes_affected.len() == 1 {
let byte = bytes_affected[0];
ComplexMapping::Single(ComplexMappingEntry {
byte,
from: self.view.get(s, byte),
to: new_values.next().unwrap(),
})
} else {
ComplexMapping::Multiple({
let mut items = bytes_affected
.iter()
.zip(new_values)
.map(|(&byte, to)| ComplexMappingEntry {
byte,
from: self.view.get(s, byte),
to,
})
.collect::<Vec<_>>();
items.retain(|item| item.from != item.to);
items
})
},
need_address_update: false,
};
data.items.push(item);
Some(ComplexJitState {
data: self.data.clone(),
id,
})
}
}
}