use crate::plumbing::CycleDetected;
use crate::plumbing::QueryDescriptor;
use crate::plumbing::QueryFunction;
use crate::plumbing::QueryStorageMassOps;
use crate::plumbing::QueryStorageOps;
use crate::plumbing::UncheckedMutQueryStorageOps;
use crate::runtime::ChangedAt;
use crate::runtime::FxIndexSet;
use crate::runtime::Revision;
use crate::runtime::Runtime;
use crate::runtime::RuntimeId;
use crate::runtime::StampedValue;
use crate::{Database, Event, EventKind, SweepStrategy};
use log::{debug, info};
use parking_lot::Mutex;
use parking_lot::RwLock;
use rustc_hash::FxHashMap;
use smallvec::SmallVec;
use std::marker::PhantomData;
use std::ops::Deref;
use std::sync::mpsc::{self, Receiver, Sender};
use std::sync::Arc;
pub type MemoizedStorage<DB, Q> = DerivedStorage<DB, Q, AlwaysMemoizeValue>;
pub type DependencyStorage<DB, Q> = DerivedStorage<DB, Q, NeverMemoizeValue>;
pub type VolatileStorage<DB, Q> = DerivedStorage<DB, Q, VolatileValue>;
pub struct DerivedStorage<DB, Q, MP>
where
Q: QueryFunction<DB>,
DB: Database,
MP: MemoizationPolicy<DB, Q>,
{
map: RwLock<FxHashMap<Q::Key, QueryState<DB, Q>>>,
policy: PhantomData<MP>,
}
impl<DB, Q, MP> std::panic::RefUnwindSafe for DerivedStorage<DB, Q, MP>
where
Q: QueryFunction<DB>,
DB: Database,
MP: MemoizationPolicy<DB, Q>,
Q::Key: std::panic::RefUnwindSafe,
Q::Value: std::panic::RefUnwindSafe,
{
}
pub trait MemoizationPolicy<DB, Q>
where
Q: QueryFunction<DB>,
DB: Database,
{
fn should_memoize_value(key: &Q::Key) -> bool;
fn memoized_value_eq(old_value: &Q::Value, new_value: &Q::Value) -> bool;
fn should_track_inputs(key: &Q::Key) -> bool;
}
pub enum AlwaysMemoizeValue {}
impl<DB, Q> MemoizationPolicy<DB, Q> for AlwaysMemoizeValue
where
Q: QueryFunction<DB>,
Q::Value: Eq,
DB: Database,
{
fn should_memoize_value(_key: &Q::Key) -> bool {
true
}
fn memoized_value_eq(old_value: &Q::Value, new_value: &Q::Value) -> bool {
old_value == new_value
}
fn should_track_inputs(_key: &Q::Key) -> bool {
true
}
}
pub enum NeverMemoizeValue {}
impl<DB, Q> MemoizationPolicy<DB, Q> for NeverMemoizeValue
where
Q: QueryFunction<DB>,
DB: Database,
{
fn should_memoize_value(_key: &Q::Key) -> bool {
false
}
fn memoized_value_eq(_old_value: &Q::Value, _new_value: &Q::Value) -> bool {
panic!("cannot reach since we never memoize")
}
fn should_track_inputs(_key: &Q::Key) -> bool {
true
}
}
pub enum VolatileValue {}
impl<DB, Q> MemoizationPolicy<DB, Q> for VolatileValue
where
Q: QueryFunction<DB>,
DB: Database,
{
fn should_memoize_value(_key: &Q::Key) -> bool {
true
}
fn memoized_value_eq(_old_value: &Q::Value, _new_value: &Q::Value) -> bool {
false
}
fn should_track_inputs(_key: &Q::Key) -> bool {
false
}
}
enum QueryState<DB, Q>
where
Q: QueryFunction<DB>,
DB: Database,
{
InProgress {
id: RuntimeId,
waiting: Mutex<SmallVec<[Sender<StampedValue<Q::Value>>; 2]>>,
},
Memoized(Memo<DB, Q>),
}
impl<DB, Q> QueryState<DB, Q>
where
Q: QueryFunction<DB>,
DB: Database,
{
fn in_progress(id: RuntimeId) -> Self {
QueryState::InProgress {
id,
waiting: Default::default(),
}
}
}
struct Memo<DB, Q>
where
Q: QueryFunction<DB>,
DB: Database,
{
value: Option<Q::Value>,
verified_at: Revision,
changed_at: Revision,
inputs: MemoInputs<DB>,
}
pub(crate) enum MemoInputs<DB: Database> {
Constant,
Tracked {
inputs: Arc<FxIndexSet<DB::QueryDescriptor>>,
},
Untracked,
}
impl<DB: Database> MemoInputs<DB> {
fn is_constant(&self) -> bool {
if let MemoInputs::Constant = self {
true
} else {
false
}
}
}
impl<DB: Database> std::fmt::Debug for MemoInputs<DB> {
fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
MemoInputs::Constant => fmt.debug_struct("Constant").finish(),
MemoInputs::Tracked { inputs } => {
fmt.debug_struct("Tracked").field("inputs", inputs).finish()
}
MemoInputs::Untracked => fmt.debug_struct("Untracked").finish(),
}
}
}
impl<DB, Q, MP> Default for DerivedStorage<DB, Q, MP>
where
Q: QueryFunction<DB>,
DB: Database,
MP: MemoizationPolicy<DB, Q>,
{
fn default() -> Self {
DerivedStorage {
map: RwLock::new(FxHashMap::default()),
policy: PhantomData,
}
}
}
enum ProbeState<V, G> {
UpToDate(Result<V, CycleDetected>),
StaleOrAbsent(G),
}
impl<DB, Q, MP> DerivedStorage<DB, Q, MP>
where
Q: QueryFunction<DB>,
DB: Database,
MP: MemoizationPolicy<DB, Q>,
{
fn read(
&self,
db: &DB,
key: &Q::Key,
descriptor: &DB::QueryDescriptor,
) -> Result<StampedValue<Q::Value>, CycleDetected> {
let runtime = db.salsa_runtime();
let revision_now = runtime.current_revision();
info!(
"{:?}({:?}): invoked at {:?}",
Q::default(),
key,
revision_now,
);
match self.probe(db, self.map.read(), runtime, revision_now, descriptor, key) {
ProbeState::UpToDate(v) => return v,
ProbeState::StaleOrAbsent(_guard) => (),
}
self.read_upgrade(db, key, descriptor, revision_now)
}
fn read_upgrade(
&self,
db: &DB,
key: &Q::Key,
descriptor: &DB::QueryDescriptor,
revision_now: Revision,
) -> Result<StampedValue<Q::Value>, CycleDetected> {
let runtime = db.salsa_runtime();
debug!(
"{:?}({:?}): read_upgrade(revision_now={:?})",
Q::default(),
key,
revision_now,
);
let mut old_memo =
match self.probe(db, self.map.write(), runtime, revision_now, descriptor, key) {
ProbeState::UpToDate(v) => return v,
ProbeState::StaleOrAbsent(mut map) => {
match map.insert(key.clone(), QueryState::in_progress(runtime.id())) {
Some(QueryState::Memoized(old_memo)) => Some(old_memo),
Some(QueryState::InProgress { .. }) => unreachable!(),
None => None,
}
}
};
let panic_guard = PanicGuard::new(&self.map, key, descriptor, runtime);
if let Some(memo) = &mut old_memo {
if let Some(value) = memo.validate_memoized_value(db, revision_now) {
info!(
"{:?}({:?}): validated old memoized value",
Q::default(),
key
);
db.salsa_event(|| Event {
runtime_id: runtime.id(),
kind: EventKind::DidValidateMemoizedValue {
descriptor: descriptor.clone(),
},
});
panic_guard.proceed(old_memo.unwrap(), &value);
return Ok(value);
}
}
let mut result = runtime.execute_query_implementation(db, descriptor, || {
info!("{:?}({:?}): executing query", Q::default(), key);
if !self.should_track_inputs(key) {
runtime.report_untracked_read();
}
Q::execute(db, key.clone())
});
assert_eq!(
runtime.current_revision(),
revision_now,
"revision altered during query execution",
);
if let Some(old_memo) = &old_memo {
if let Some(old_value) = &old_memo.value {
if MP::memoized_value_eq(&old_value, &result.value) {
debug!(
"read_upgrade({:?}({:?})): value is equal, back-dating to {:?}",
Q::default(),
key,
old_memo.changed_at,
);
assert!(old_memo.changed_at <= result.changed_at.revision);
result.changed_at.revision = old_memo.changed_at;
}
}
}
let new_value = StampedValue {
value: result.value,
changed_at: result.changed_at,
};
let value = if self.should_memoize_value(key) {
Some(new_value.value.clone())
} else {
None
};
debug!(
"read_upgrade({:?}({:?})): result.changed_at={:?}, result.subqueries = {:#?}",
Q::default(),
key,
result.changed_at,
result.subqueries,
);
let inputs = match result.subqueries {
None => MemoInputs::Untracked,
Some(descriptors) => {
if descriptors.is_empty() || result.changed_at.is_constant {
MemoInputs::Constant
} else {
MemoInputs::Tracked {
inputs: Arc::new(descriptors),
}
}
}
};
panic_guard.proceed(
Memo {
value,
changed_at: result.changed_at.revision,
verified_at: revision_now,
inputs,
},
&new_value,
);
Ok(new_value)
}
fn probe<MapGuard>(
&self,
db: &DB,
map: MapGuard,
runtime: &Runtime<DB>,
revision_now: Revision,
descriptor: &DB::QueryDescriptor,
key: &Q::Key,
) -> ProbeState<StampedValue<Q::Value>, MapGuard>
where
MapGuard: Deref<Target = FxHashMap<Q::Key, QueryState<DB, Q>>>,
{
match map.get(key) {
Some(QueryState::InProgress { id, waiting }) => {
let other_id = *id;
return match self
.register_with_in_progress_thread(runtime, descriptor, other_id, waiting)
{
Ok(rx) => {
std::mem::drop(map);
db.salsa_event(|| Event {
runtime_id: db.salsa_runtime().id(),
kind: EventKind::WillBlockOn {
other_runtime_id: other_id,
descriptor: descriptor.clone(),
},
});
let value = rx.recv().unwrap_or_else(|_| db.on_propagated_panic());
ProbeState::UpToDate(Ok(value))
}
Err(CycleDetected) => ProbeState::UpToDate(Err(CycleDetected)),
};
}
Some(QueryState::Memoized(memo)) => {
debug!("{:?}({:?}): found memoized value", Q::default(), key);
if let Some(value) = memo.probe_memoized_value(revision_now) {
info!(
"{:?}({:?}): returning memoized value changed at {:?}",
Q::default(),
key,
value.changed_at
);
return ProbeState::UpToDate(Ok(value));
}
}
None => {}
}
ProbeState::StaleOrAbsent(map)
}
fn register_with_in_progress_thread(
&self,
runtime: &Runtime<DB>,
descriptor: &DB::QueryDescriptor,
other_id: RuntimeId,
waiting: &Mutex<SmallVec<[Sender<StampedValue<Q::Value>>; 2]>>,
) -> Result<Receiver<StampedValue<Q::Value>>, CycleDetected> {
if other_id == runtime.id() {
return Err(CycleDetected);
} else {
if !runtime.try_block_on(descriptor, other_id) {
return Err(CycleDetected);
}
let (tx, rx) = mpsc::channel();
waiting.lock().push(tx);
Ok(rx)
}
}
fn should_memoize_value(&self, key: &Q::Key) -> bool {
MP::should_memoize_value(key)
}
fn should_track_inputs(&self, key: &Q::Key) -> bool {
MP::should_track_inputs(key)
}
}
struct PanicGuard<'db, DB, Q>
where
DB: Database,
Q: QueryFunction<DB>,
{
descriptor: &'db DB::QueryDescriptor,
key: &'db Q::Key,
map: &'db RwLock<FxHashMap<Q::Key, QueryState<DB, Q>>>,
runtime: &'db Runtime<DB>,
}
impl<'db, DB, Q> PanicGuard<'db, DB, Q>
where
DB: Database + 'db,
Q: QueryFunction<DB>,
{
fn new(
map: &'db RwLock<FxHashMap<Q::Key, QueryState<DB, Q>>>,
key: &'db Q::Key,
descriptor: &'db DB::QueryDescriptor,
runtime: &'db Runtime<DB>,
) -> Self {
Self {
descriptor,
key,
map,
runtime,
}
}
fn proceed(self, memo: Memo<DB, Q>, new_value: &StampedValue<Q::Value>) {
self.overwrite_placeholder(Some(memo), Some(new_value));
std::mem::forget(self)
}
fn overwrite_placeholder(
&self,
memo: Option<Memo<DB, Q>>,
new_value: Option<&StampedValue<Q::Value>>,
) {
let mut write = self.map.write();
let old_value = match memo {
Some(memo) => write.insert(self.key.clone(), QueryState::Memoized(memo)),
None => write.remove(self.key),
};
match old_value {
Some(QueryState::InProgress { id, waiting }) => {
assert_eq!(id, self.runtime.id());
self.runtime
.unblock_queries_blocked_on_self(self.descriptor);
match new_value {
Some(new_value) => {
for tx in waiting.into_inner() {
tx.send(new_value.clone()).unwrap()
}
}
None => std::mem::drop(waiting),
}
}
_ => panic!(
"\
Unexpected panic during query evaluation, aborting the process.
Please report this bug to https://github.com/salsa-rs/salsa/issues."
),
}
}
}
impl<'db, DB, Q> Drop for PanicGuard<'db, DB, Q>
where
DB: Database + 'db,
Q: QueryFunction<DB>,
{
fn drop(&mut self) {
if std::thread::panicking() {
self.overwrite_placeholder(None, None)
} else {
panic!(".forget() was not called")
}
}
}
impl<DB, Q, MP> QueryStorageOps<DB, Q> for DerivedStorage<DB, Q, MP>
where
Q: QueryFunction<DB>,
DB: Database,
MP: MemoizationPolicy<DB, Q>,
{
fn try_fetch(
&self,
db: &DB,
key: &Q::Key,
descriptor: &DB::QueryDescriptor,
) -> Result<Q::Value, CycleDetected> {
let StampedValue { value, changed_at } = self.read(db, key, &descriptor)?;
db.salsa_runtime().report_query_read(descriptor, changed_at);
Ok(value)
}
fn maybe_changed_since(
&self,
db: &DB,
revision: Revision,
key: &Q::Key,
descriptor: &DB::QueryDescriptor,
) -> bool {
let runtime = db.salsa_runtime();
let revision_now = runtime.current_revision();
debug!(
"maybe_changed_since({:?}({:?})) called with revision={:?}, revision_now={:?}",
Q::default(),
key,
revision,
revision_now,
);
let map = self.map.read();
let memo = match map.get(key) {
None => {
debug!(
"maybe_changed_since({:?}({:?}): no value",
Q::default(),
key,
);
return true;
}
Some(QueryState::InProgress { id, waiting }) => {
let other_id = *id;
debug!(
"maybe_changed_since({:?}({:?}): blocking on thread `{:?}`",
Q::default(),
key,
other_id,
);
match self.register_with_in_progress_thread(runtime, descriptor, other_id, waiting)
{
Ok(rx) => {
std::mem::drop(map);
let value = rx.recv().unwrap_or_else(|_| db.on_propagated_panic());
return value.changed_at.changed_since(revision);
}
Err(CycleDetected) => return true,
}
}
Some(QueryState::Memoized(memo)) => memo,
};
if memo.verified_at == revision_now {
debug!(
"maybe_changed_since({:?}({:?}): {:?} since up-to-date memo that changed at {:?}",
Q::default(),
key,
memo.changed_at > revision,
memo.changed_at,
);
return memo.changed_at > revision;
}
let inputs = match &memo.inputs {
MemoInputs::Untracked => {
debug!(
"maybe_changed_since({:?}({:?}): true since untracked inputs",
Q::default(),
key,
);
return true;
}
MemoInputs::Constant => None,
MemoInputs::Tracked { inputs } => {
assert!(inputs.len() > 0);
if memo.value.is_some() {
std::mem::drop(map);
return match self.read_upgrade(db, key, descriptor, revision_now) {
Ok(v) => {
debug!(
"maybe_changed_since({:?}({:?}): {:?} since (recomputed) value changed at {:?}",
Q::default(),
key,
v.changed_at.changed_since(revision),
v.changed_at,
);
v.changed_at.changed_since(revision)
}
Err(CycleDetected) => true,
};
}
Some(inputs.clone())
}
};
std::mem::drop(map);
let maybe_changed = inputs
.iter()
.flat_map(|inputs| inputs.iter())
.filter(|input| input.maybe_changed_since(db, revision))
.inspect(|input| {
debug!(
"{:?}({:?}): input `{:?}` may have changed",
Q::default(),
key,
input
)
})
.next()
.is_some();
{
let mut map = self.map.write();
match map.get_mut(key) {
Some(QueryState::Memoized(memo)) => {
if memo.verified_at == revision_now {
} else if maybe_changed {
map.remove(key);
} else {
memo.verified_at = revision_now;
}
}
Some(QueryState::InProgress { .. }) => {
}
None => {
}
}
}
maybe_changed
}
fn is_constant(&self, _db: &DB, key: &Q::Key) -> bool {
let map_read = self.map.read();
match map_read.get(key) {
None => false,
Some(QueryState::InProgress { .. }) => panic!("query in progress"),
Some(QueryState::Memoized(memo)) => memo.inputs.is_constant(),
}
}
fn keys<C>(&self, _db: &DB) -> C
where
C: std::iter::FromIterator<Q::Key>,
{
let map = self.map.read();
map.keys().cloned().collect()
}
}
impl<DB, Q, MP> QueryStorageMassOps<DB> for DerivedStorage<DB, Q, MP>
where
Q: QueryFunction<DB>,
DB: Database,
MP: MemoizationPolicy<DB, Q>,
{
fn sweep(&self, db: &DB, strategy: SweepStrategy) {
let mut map_write = self.map.write();
let revision_now = db.salsa_runtime().current_revision();
map_write.retain(|key, query_state| {
match query_state {
QueryState::InProgress { .. } => {
debug!("sweep({:?}({:?})): in-progress", Q::default(), key);
true
}
QueryState::Memoized(memo) => {
debug!(
"sweep({:?}({:?})): last verified at {:?}, current revision {:?}",
Q::default(),
key,
memo.verified_at,
revision_now
);
assert!(memo.verified_at <= revision_now);
if !strategy.keep_values {
memo.value = None;
}
memo.verified_at == revision_now
}
}
});
}
}
impl<DB, Q, MP> UncheckedMutQueryStorageOps<DB, Q> for DerivedStorage<DB, Q, MP>
where
Q: QueryFunction<DB>,
DB: Database,
MP: MemoizationPolicy<DB, Q>,
{
fn set_unchecked(&self, db: &DB, key: &Q::Key, value: Q::Value) {
let key = key.clone();
let mut map_write = self.map.write();
let current_revision = db.salsa_runtime().current_revision();
map_write.insert(
key,
QueryState::Memoized(Memo {
value: Some(value),
changed_at: current_revision,
verified_at: current_revision,
inputs: MemoInputs::Tracked {
inputs: Default::default(),
},
}),
);
}
}
impl<DB, Q> Memo<DB, Q>
where
Q: QueryFunction<DB>,
DB: Database,
{
fn validate_memoized_value(
&mut self,
db: &DB,
revision_now: Revision,
) -> Option<StampedValue<Q::Value>> {
let value = self.value.as_ref()?;
assert!(self.verified_at != revision_now);
let verified_at = self.verified_at;
debug!(
"validate_memoized_value({:?}): verified_at={:#?}",
Q::default(),
self.inputs,
);
let is_constant = match &mut self.inputs {
MemoInputs::Untracked { .. } => {
return None;
}
MemoInputs::Constant => true,
MemoInputs::Tracked { inputs } => {
let changed_input = inputs
.iter()
.filter(|input| input.maybe_changed_since(db, verified_at))
.next();
if let Some(input) = changed_input {
debug!(
"{:?}::validate_memoized_value: `{:?}` may have changed",
Q::default(),
input
);
return None;
}
false
}
};
self.verified_at = revision_now;
Some(StampedValue {
changed_at: ChangedAt {
is_constant,
revision: self.changed_at,
},
value: value.clone(),
})
}
fn probe_memoized_value(&self, revision_now: Revision) -> Option<StampedValue<Q::Value>> {
let value = self.value.as_ref()?;
debug!(
"probe_memoized_value(verified_at={:?}, changed_at={:?})",
self.verified_at, self.changed_at,
);
if self.verified_at == revision_now {
let is_constant = match self.inputs {
MemoInputs::Constant => true,
_ => false,
};
return Some(StampedValue {
changed_at: ChangedAt {
is_constant,
revision: self.changed_at,
},
value: value.clone(),
});
}
None
}
}