use crate::{Database, Event, EventKind, SweepStrategy};
use lock_api::{RawRwLock, RawRwLockRecursive};
use log::debug;
use parking_lot::{Mutex, RwLock};
use rustc_hash::{FxHashMap, FxHasher};
use smallvec::SmallVec;
use std::cell::RefCell;
use std::fmt::Write;
use std::hash::BuildHasherDefault;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
pub(crate) type FxIndexSet<K> = indexmap::IndexSet<K, BuildHasherDefault<FxHasher>>;
pub struct Runtime<DB: Database> {
id: RuntimeId,
revision_guard: Option<RevisionGuard<DB>>,
local_state: RefCell<LocalState<DB>>,
shared_state: Arc<SharedState<DB>>,
}
impl<DB> std::panic::RefUnwindSafe for Runtime<DB>
where
DB: Database,
DB::DatabaseStorage: std::panic::RefUnwindSafe,
{
}
impl<DB> Default for Runtime<DB>
where
DB: Database,
{
fn default() -> Self {
Runtime {
id: RuntimeId { counter: 0 },
revision_guard: None,
shared_state: Default::default(),
local_state: Default::default(),
}
}
}
impl<DB> std::fmt::Debug for Runtime<DB>
where
DB: Database,
{
fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
fmt.debug_struct("Runtime")
.field("id", &self.id())
.field("forked", &self.revision_guard.is_some())
.field("shared_state", &self.shared_state)
.finish()
}
}
impl<DB> Runtime<DB>
where
DB: Database,
{
pub fn new() -> Self {
Self::default()
}
pub fn storage(&self) -> &DB::DatabaseStorage {
&self.shared_state.storage
}
pub fn snapshot(&self, from_db: &DB) -> Self {
assert!(
Arc::ptr_eq(&self.shared_state, &from_db.salsa_runtime().shared_state),
"invoked `snapshot` with a non-matching database"
);
if self.local_state.borrow().query_in_progress() {
panic!("it is not legal to `snapshot` during a query (see salsa-rs/salsa#80)");
}
let revision_guard = RevisionGuard::new(&self.shared_state);
let id = RuntimeId {
counter: self.shared_state.next_id.fetch_add(1, Ordering::SeqCst),
};
Runtime {
id,
revision_guard: Some(revision_guard),
shared_state: self.shared_state.clone(),
local_state: Default::default(),
}
}
pub fn next_revision(&self) {
self.with_incremented_revision(|_| ());
}
pub fn sweep_all(&self, db: &DB, strategy: SweepStrategy) {
db.for_each_query(|query_storage| query_storage.sweep(db, strategy));
}
#[inline]
pub fn id(&self) -> RuntimeId {
self.id
}
pub fn active_query(&self) -> Option<DB::QueryDescriptor> {
self.local_state
.borrow()
.query_stack
.last()
.map(|active_query| active_query.descriptor.clone())
}
#[inline]
pub(crate) fn current_revision(&self) -> Revision {
Revision {
generation: self.shared_state.revision.load(Ordering::SeqCst) as u64,
}
}
#[inline]
fn pending_revision(&self) -> Revision {
Revision {
generation: self.shared_state.pending_revision.load(Ordering::SeqCst) as u64,
}
}
#[inline]
pub fn is_current_revision_canceled(&self) -> bool {
let current_revision = self.current_revision();
let pending_revision = self.pending_revision();
debug!(
"is_current_revision_canceled: current_revision={:?}, pending_revision={:?}",
current_revision, pending_revision
);
if pending_revision > current_revision {
self.report_untracked_read();
true
} else {
assert_eq!(pending_revision, current_revision);
self.report_anon_read(pending_revision);
false
}
}
pub(crate) fn with_incremented_revision<R>(&self, op: impl FnOnce(Revision) -> R) -> R {
log::debug!("increment_revision()");
if !self.permits_increment() {
panic!("increment_revision invoked during a query computation");
}
let current_revision = self
.shared_state
.pending_revision
.fetch_add(1, Ordering::SeqCst);
assert!(current_revision != usize::max_value(), "revision overflow");
let _lock = self.shared_state.query_lock.write();
let old_revision = self.shared_state.revision.fetch_add(1, Ordering::SeqCst);
assert_eq!(current_revision, old_revision);
let new_revision = Revision {
generation: (current_revision + 1) as u64,
};
debug!("increment_revision: incremented to {:?}", new_revision);
op(new_revision)
}
pub(crate) fn permits_increment(&self) -> bool {
self.revision_guard.is_none() && !self.local_state.borrow().query_in_progress()
}
pub(crate) fn execute_query_implementation<V>(
&self,
db: &DB,
descriptor: &DB::QueryDescriptor,
execute: impl FnOnce() -> V,
) -> ComputedQueryResult<DB, V> {
debug!("{:?}: execute_query_implementation invoked", descriptor);
db.salsa_event(|| Event {
runtime_id: db.salsa_runtime().id(),
kind: EventKind::WillExecute {
descriptor: descriptor.clone(),
},
});
let push_len = {
let mut local_state = self.local_state.borrow_mut();
local_state
.query_stack
.push(ActiveQuery::new(descriptor.clone()));
local_state.query_stack.len()
};
let value = execute();
let ActiveQuery {
subqueries,
changed_at,
..
} = {
let mut local_state = self.local_state.borrow_mut();
assert_eq!(local_state.query_stack.len(), push_len);
local_state.query_stack.pop().unwrap()
};
ComputedQueryResult {
value,
changed_at,
subqueries,
}
}
pub(crate) fn report_query_read(
&self,
descriptor: &DB::QueryDescriptor,
changed_at: ChangedAt,
) {
if let Some(top_query) = self.local_state.borrow_mut().query_stack.last_mut() {
top_query.add_read(descriptor, changed_at);
}
}
pub(crate) fn report_untracked_read(&self) {
if let Some(top_query) = self.local_state.borrow_mut().query_stack.last_mut() {
top_query.add_untracked_read(self.current_revision());
}
}
fn report_anon_read(&self, revision: Revision) {
if let Some(top_query) = self.local_state.borrow_mut().query_stack.last_mut() {
top_query.add_anon_read(revision);
}
}
pub(crate) fn report_unexpected_cycle(&self, descriptor: DB::QueryDescriptor) -> ! {
debug!("report_unexpected_cycle(descriptor={:?})", descriptor);
let local_state = self.local_state.borrow();
let LocalState { query_stack, .. } = &*local_state;
let start_index = (0..query_stack.len())
.rev()
.filter(|&i| query_stack[i].descriptor == descriptor)
.next()
.unwrap();
let mut message = format!("Internal error, cycle detected:\n");
for active_query in &query_stack[start_index..] {
writeln!(message, "- {:?}\n", active_query.descriptor).unwrap();
}
panic!(message)
}
pub(crate) fn try_block_on(
&self,
descriptor: &DB::QueryDescriptor,
other_id: RuntimeId,
) -> bool {
self.shared_state
.dependency_graph
.lock()
.add_edge(self.id(), descriptor, other_id)
}
pub(crate) fn unblock_queries_blocked_on_self(&self, descriptor: &DB::QueryDescriptor) {
self.shared_state
.dependency_graph
.lock()
.remove_edge(descriptor, self.id())
}
}
struct SharedState<DB: Database> {
storage: DB::DatabaseStorage,
next_id: AtomicUsize,
query_lock: RwLock<()>,
revision: AtomicUsize,
pending_revision: AtomicUsize,
dependency_graph: Mutex<DependencyGraph<DB>>,
}
impl<DB: Database> Default for SharedState<DB> {
fn default() -> Self {
SharedState {
next_id: AtomicUsize::new(1),
storage: Default::default(),
query_lock: Default::default(),
revision: Default::default(),
pending_revision: Default::default(),
dependency_graph: Default::default(),
}
}
}
impl<DB> std::fmt::Debug for SharedState<DB>
where
DB: Database,
{
fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let query_lock = if self.query_lock.try_write().is_some() {
"<unlocked>"
} else if self.query_lock.try_read().is_some() {
"<rlocked>"
} else {
"<wlocked>"
};
fmt.debug_struct("SharedState")
.field("query_lock", &query_lock)
.field("revision", &self.revision)
.field("pending_revision", &self.pending_revision)
.finish()
}
}
struct LocalState<DB: Database> {
query_stack: Vec<ActiveQuery<DB>>,
}
impl<DB: Database> Default for LocalState<DB> {
fn default() -> Self {
LocalState {
query_stack: Default::default(),
}
}
}
impl<DB: Database> LocalState<DB> {
fn query_in_progress(&self) -> bool {
!self.query_stack.is_empty()
}
}
struct ActiveQuery<DB: Database> {
descriptor: DB::QueryDescriptor,
changed_at: ChangedAt,
subqueries: Option<FxIndexSet<DB::QueryDescriptor>>,
}
pub(crate) struct ComputedQueryResult<DB: Database, V> {
pub(crate) value: V,
pub(crate) changed_at: ChangedAt,
pub(crate) subqueries: Option<FxIndexSet<DB::QueryDescriptor>>,
}
impl<DB: Database> ActiveQuery<DB> {
fn new(descriptor: DB::QueryDescriptor) -> Self {
ActiveQuery {
descriptor,
changed_at: ChangedAt {
is_constant: true,
revision: Revision::ZERO,
},
subqueries: Some(FxIndexSet::default()),
}
}
fn add_read(&mut self, subquery: &DB::QueryDescriptor, changed_at: ChangedAt) {
let ChangedAt {
is_constant,
revision,
} = changed_at;
if let Some(set) = &mut self.subqueries {
set.insert(subquery.clone());
}
self.changed_at.is_constant &= is_constant;
self.changed_at.revision = self.changed_at.revision.max(revision);
}
fn add_untracked_read(&mut self, changed_at: Revision) {
self.subqueries = None;
self.changed_at.is_constant = false;
self.changed_at.revision = changed_at;
}
fn add_anon_read(&mut self, changed_at: Revision) {
self.changed_at.revision = self.changed_at.revision.max(changed_at);
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct RuntimeId {
counter: usize,
}
#[derive(Copy, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct Revision {
generation: u64,
}
impl Revision {
pub(crate) const ZERO: Self = Revision { generation: 0 };
fn next(self) -> Revision {
Revision {
generation: self.generation + 1,
}
}
fn as_usize(self) -> usize {
assert!(self.generation < (std::usize::MAX as u64));
self.generation as usize
}
}
impl std::fmt::Debug for Revision {
fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(fmt, "R{}", self.generation)
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub struct ChangedAt {
pub(crate) is_constant: bool,
pub(crate) revision: Revision,
}
impl ChangedAt {
pub(crate) fn changed_since(self, revision: Revision) -> bool {
self.revision > revision
}
}
#[derive(Clone, Debug)]
pub(crate) struct StampedValue<V> {
pub(crate) value: V,
pub(crate) changed_at: ChangedAt,
}
struct DependencyGraph<DB: Database> {
edges: FxHashMap<RuntimeId, RuntimeId>,
labels: FxHashMap<DB::QueryDescriptor, SmallVec<[RuntimeId; 4]>>,
}
impl<DB: Database> Default for DependencyGraph<DB> {
fn default() -> Self {
DependencyGraph {
edges: Default::default(),
labels: Default::default(),
}
}
}
impl<DB: Database> DependencyGraph<DB> {
fn add_edge(
&mut self,
from_id: RuntimeId,
descriptor: &DB::QueryDescriptor,
to_id: RuntimeId,
) -> bool {
assert_ne!(from_id, to_id);
debug_assert!(!self.edges.contains_key(&from_id));
let mut p = to_id;
while let Some(&q) = self.edges.get(&p) {
if q == from_id {
return false;
}
p = q;
}
self.edges.insert(from_id, to_id);
self.labels
.entry(descriptor.clone())
.or_insert(SmallVec::default())
.push(from_id);
true
}
fn remove_edge(&mut self, descriptor: &DB::QueryDescriptor, to_id: RuntimeId) {
let vec = self
.labels
.remove(descriptor)
.unwrap_or(SmallVec::default());
for from_id in &vec {
let to_id1 = self.edges.remove(from_id);
assert_eq!(Some(to_id), to_id1);
}
}
}
struct RevisionGuard<DB: Database> {
shared_state: Arc<SharedState<DB>>,
}
impl<DB> RevisionGuard<DB>
where
DB: Database,
{
fn new(shared_state: &Arc<SharedState<DB>>) -> Self {
unsafe {
shared_state.query_lock.raw().lock_shared_recursive();
}
Self {
shared_state: shared_state.clone(),
}
}
}
impl<DB> Drop for RevisionGuard<DB>
where
DB: Database,
{
fn drop(&mut self) {
unsafe {
self.shared_state.query_lock.raw().unlock_shared();
}
}
}