Skip to content

Commit 6ecfd9f

Browse files
committed
Use the extension mechanism for AnalysisCtxt
1 parent 5463189 commit 6ecfd9f

4 files changed

Lines changed: 47 additions & 66 deletions

File tree

‎src/atomic_context.rs‎

Lines changed: 2 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -298,15 +298,7 @@ impl<'tcx> AnalysisCtxt<'tcx> {
298298
}
299299

300300
pub struct AtomicContext<'tcx> {
301-
cx: AnalysisCtxt<'tcx>,
302-
}
303-
304-
impl<'tcx> AtomicContext<'tcx> {
305-
pub fn new(tcx: TyCtxt<'tcx>) -> Self {
306-
Self {
307-
cx: AnalysisCtxt::new(tcx),
308-
}
309-
}
301+
pub cx: &'tcx AnalysisCtxt<'tcx>,
310302
}
311303

312304
impl_lint_pass!(AtomicContext<'_> => [ATOMIC_CONTEXT]);
@@ -440,7 +432,7 @@ impl<'tcx> LateLintPass<'tcx> for AtomicContext<'tcx> {
440432
.tcx
441433
.erase_and_anonymize_regions(GenericArgs::identity_for_item(self.cx.tcx, def_id));
442434
let instance = Instance::new_raw(def_id.into(), identity);
443-
let poly_instance = TypingEnv::post_analysis(*self.cx, def_id).as_query_input(instance);
435+
let poly_instance = TypingEnv::post_analysis(self.cx.tcx, def_id).as_query_input(instance);
444436
let _ = self.cx.instance_adjustment(poly_instance);
445437
let _ = self.cx.instance_expectation(poly_instance);
446438
let _ = self.cx.instance_check(poly_instance);

‎src/ctxt.rs‎

Lines changed: 23 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -3,11 +3,11 @@
33
// SPDX-License-Identifier: MIT OR Apache-2.0
44

55
use std::any::{Any, TypeId};
6-
use std::cell::RefCell;
76
use std::sync::Arc;
87

98
use rusqlite::{Connection, OptionalExtension};
109
use rustc_data_structures::fx::FxHashMap;
10+
use rustc_data_structures::sync::{DynSend, DynSync, MTLock, RwLock};
1111
use rustc_hir::def_id::{CrateNum, LOCAL_CRATE};
1212
use rustc_middle::ty::TyCtxt;
1313
use rustc_serialize::{Decodable, Encodable};
@@ -18,8 +18,8 @@ use crate::preempt_count::UseSite;
1818
pub(crate) trait Query: 'static {
1919
const NAME: &'static str;
2020

21-
type Key<'tcx>;
22-
type Value<'tcx>;
21+
type Key<'tcx>: DynSend + DynSync;
22+
type Value<'tcx>: DynSend + DynSync;
2323
}
2424

2525
pub(crate) trait QueryValueDecodable: Query {
@@ -50,13 +50,17 @@ pub(crate) trait PersistentQuery: QueryValueDecodable {
5050

5151
pub struct AnalysisCtxt<'tcx> {
5252
pub tcx: TyCtxt<'tcx>,
53-
pub local_conn: Connection,
54-
pub sql_conn: RefCell<FxHashMap<CrateNum, Option<Arc<Connection>>>>,
53+
pub local_conn: MTLock<Connection>,
54+
pub sql_conn: RwLock<FxHashMap<CrateNum, Option<Arc<MTLock<Connection>>>>>,
5555

56-
pub call_stack: RefCell<Vec<UseSite<'tcx>>>,
57-
pub query_cache: RefCell<FxHashMap<TypeId, Arc<dyn Any>>>,
56+
pub call_stack: RwLock<Vec<UseSite<'tcx>>>,
57+
pub query_cache: RwLock<FxHashMap<TypeId, Arc<dyn Any + DynSend + DynSync>>>,
5858
}
5959

60+
// Everything in `AnalysisCtxt` is either `DynSend/DynSync` or `Send/Sync`, but since there're no relation between two right now compiler cannot infer this.
61+
unsafe impl<'tcx> DynSend for AnalysisCtxt<'tcx> {}
62+
unsafe impl<'tcx> DynSync for AnalysisCtxt<'tcx> {}
63+
6064
impl<'tcx> std::ops::Deref for AnalysisCtxt<'tcx> {
6165
type Target = TyCtxt<'tcx>;
6266

@@ -105,7 +109,7 @@ const SCHEMA_VERSION: u32 = 1;
105109

106110
impl Drop for AnalysisCtxt<'_> {
107111
fn drop(&mut self) {
108-
self.local_conn.execute("commit", ()).unwrap();
112+
self.local_conn.lock().execute("commit", ()).unwrap();
109113
}
110114
}
111115

@@ -127,26 +131,26 @@ impl ArcDowncast for Arc<dyn Any> {
127131
impl<'tcx> AnalysisCtxt<'tcx> {
128132
pub(crate) fn query_cache<Q: Query>(
129133
&self,
130-
) -> Arc<RefCell<FxHashMap<Q::Key<'tcx>, Q::Value<'tcx>>>> {
134+
) -> Arc<RwLock<FxHashMap<Q::Key<'tcx>, Q::Value<'tcx>>>> {
131135
let key = TypeId::of::<Q>();
132136
let mut guard = self.query_cache.borrow_mut();
133-
let cache = guard
137+
let cache = (guard
134138
.entry(key)
135139
.or_insert_with(|| {
136-
let cache = Arc::new(RefCell::new(
140+
let cache = Arc::new(RwLock::new(
137141
FxHashMap::<Q::Key<'static>, Q::Value<'static>>::default(),
138142
));
139143
cache
140144
})
141-
.clone()
142-
.downcast::<RefCell<FxHashMap<Q::Key<'static>, Q::Value<'static>>>>()
145+
.clone() as Arc<dyn Any>)
146+
.downcast::<RwLock<FxHashMap<Q::Key<'static>, Q::Value<'static>>>>()
143147
.unwrap();
144148
// Everything stored inside query_cache is conceptually `'tcx`, but due to limitation
145149
// of `Any` we hack around the lifetime.
146150
unsafe { std::mem::transmute(cache) }
147151
}
148152

149-
pub(crate) fn sql_connection(&self, cnum: CrateNum) -> Option<Arc<Connection>> {
153+
pub(crate) fn sql_connection(&self, cnum: CrateNum) -> Option<Arc<MTLock<Connection>>> {
150154
if let Some(v) = self.sql_conn.borrow().get(&cnum) {
151155
return v.clone();
152156
}
@@ -187,7 +191,7 @@ impl<'tcx> AnalysisCtxt<'tcx> {
187191
);
188192
}
189193

190-
result = Some(Arc::new(conn));
194+
result = Some(Arc::new(MTLock::new(conn)));
191195
break;
192196
}
193197
}
@@ -205,6 +209,7 @@ impl<'tcx> AnalysisCtxt<'tcx> {
205209

206210
pub(crate) fn sql_create_table<Q: Query>(&self) {
207211
self.local_conn
212+
.lock()
208213
.execute_batch(&format!(
209214
"CREATE TABLE {} (key BLOB PRIMARY KEY, value BLOB);",
210215
Q::NAME
@@ -225,6 +230,7 @@ impl<'tcx> AnalysisCtxt<'tcx> {
225230

226231
let value_encoded: Vec<u8> = self
227232
.sql_connection(cnum)?
233+
.lock()
228234
.query_row(
229235
&format!("SELECT value FROM {} WHERE key = ?", Q::NAME),
230236
rusqlite::params![encoded],
@@ -265,6 +271,7 @@ impl<'tcx> AnalysisCtxt<'tcx> {
265271
let value_encoded = encode_ctx.finish();
266272

267273
self.local_conn
274+
.lock()
268275
.execute(
269276
&format!(
270277
"INSERT OR REPLACE INTO {} (key, value) VALUES (?, ?)",
@@ -309,7 +316,7 @@ impl<'tcx> AnalysisCtxt<'tcx> {
309316

310317
let ret = Self {
311318
tcx,
312-
local_conn: conn,
319+
local_conn: MTLock::new(conn),
313320
sql_conn: Default::default(),
314321
call_stack: Default::default(),
315322
query_cache: Default::default(),

‎src/main.rs‎

Lines changed: 15 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,8 @@ use rustc_session::EarlyDiagCtxt;
4848
use rustc_session::config::ErrorOutputType;
4949
use std::sync::atomic::Ordering;
5050

51+
use crate::ctxt::AnalysisCtxt;
52+
5153
#[macro_use]
5254
mod ctxt;
5355

@@ -78,14 +80,16 @@ impl Callbacks for MyCallbacks {
7880
config.override_queries = Some(|_, provider| {
7981
// Calling `optimized_mir` will steal the result of query `mir_drops_elaborated_and_const_checked`,
8082
// so hijack `optimized_mir` to run `analysis_mir` first.
81-
hook_query!(provider.optimized_mir => |tcx, def_id, original| {
83+
hook_query!(provider.optimized_mir => |tcx, local_def_id, original| {
84+
let def_id = local_def_id.to_def_id();
8285
// Skip `analysis_mir` call if this is a constructor, since it will be delegated back to
8386
// `optimized_mir` for building ADT constructor shim.
84-
if !tcx.is_constructor(def_id.to_def_id()) {
85-
crate::mir::local_analysis_mir(tcx, def_id);
87+
if !tcx.is_constructor(def_id) {
88+
let cx = crate::driver::cx::<MyCallbacks>(tcx);
89+
let _ = cx.analysis_mir(def_id);
8690
}
8791

88-
original(tcx, def_id)
92+
original(tcx, local_def_id)
8993
});
9094
});
9195
config.register_lints = Some(Box::new(move |_, lint_store| {
@@ -94,16 +98,20 @@ impl Callbacks for MyCallbacks {
9498
lint_store.register_lints(&[&atomic_context::ATOMIC_CONTEXT]);
9599
// lint_store
96100
// .register_late_pass(|_| Box::new(infallible_allocation::InfallibleAllocation));
97-
lint_store.register_late_pass(|tcx| Box::new(atomic_context::AtomicContext::new(tcx)));
101+
lint_store.register_late_pass(|tcx| {
102+
Box::new(atomic_context::AtomicContext {
103+
cx: driver::cx::<MyCallbacks>(tcx),
104+
})
105+
});
98106
}));
99107
}
100108
}
101109

102110
impl driver::CallbacksExt for MyCallbacks {
103-
type ExtCtxt<'tcx> = TyCtxt<'tcx>;
111+
type ExtCtxt<'tcx> = AnalysisCtxt<'tcx>;
104112

105113
fn ext_cx<'tcx>(&mut self, tcx: TyCtxt<'tcx>) -> Self::ExtCtxt<'tcx> {
106-
tcx
114+
AnalysisCtxt::new(tcx)
107115
}
108116
}
109117

‎src/mir.rs‎

Lines changed: 7 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -6,11 +6,6 @@ pub mod drop_shim;
66
pub mod elaborate_drop;
77
pub mod patch;
88

9-
use std::sync::atomic::AtomicPtr;
10-
use std::sync::atomic::Ordering;
11-
use std::sync::{LazyLock, Mutex};
12-
13-
use rustc_data_structures::fx::FxHashMap;
149
use rustc_hir::{self as hir, def::DefKind};
1510
use rustc_middle::mir::CallSource;
1611
use rustc_middle::mir::{
@@ -24,38 +19,17 @@ use rustc_span::{DUMMY_SP, source_map::Spanned, sym};
2419
use crate::ctxt::AnalysisCtxt;
2520
use crate::ctxt::PersistentQuery;
2621

27-
// HACK: we can't add new queries to `TyCtxt` without changing rustc code, so
28-
// use this as a "poor man's query" for now.
29-
//
30-
// `AnalysisCtxt` has its own cache but we can't use it as we've only got `TyCtxt`
31-
// so far but not `AnalysisCtxt`.
32-
static MIR_CACHE: LazyLock<Mutex<FxHashMap<LocalDefId, AtomicPtr<()>>>> =
33-
LazyLock::new(|| Mutex::new(FxHashMap::default()));
34-
35-
pub fn local_analysis_mir<'tcx>(tcx: TyCtxt<'tcx>, did: LocalDefId) -> &'tcx Body<'tcx> {
36-
if tcx.is_constructor(did.to_def_id()) {
37-
return tcx.optimized_mir(did.to_def_id());
22+
pub fn local_analysis_mir<'tcx>(cx: &AnalysisCtxt<'tcx>, did: LocalDefId) -> &'tcx Body<'tcx> {
23+
if cx.is_constructor(did.to_def_id()) {
24+
return cx.optimized_mir(did.to_def_id());
3825
}
3926

40-
{
41-
let cache = MIR_CACHE.lock().unwrap();
42-
if let Some(body) = cache.get(&did) {
43-
return unsafe { &*body.load(Ordering::Relaxed).cast() };
44-
}
45-
}
46-
47-
let body = tcx
27+
let body = cx
4828
.mir_drops_elaborated_and_const_checked(did)
4929
.borrow()
5030
.clone();
51-
let body = remap_mir_for_const_eval_select(tcx, body, hir::Constness::NotConst);
52-
let body = tcx.arena.alloc(body);
53-
54-
{
55-
let mut cache = MIR_CACHE.lock().unwrap();
56-
cache.insert(did, AtomicPtr::new(body as *const _ as *mut _));
57-
}
58-
body
31+
let body = remap_mir_for_const_eval_select(cx.tcx, body, hir::Constness::NotConst);
32+
cx.arena.alloc(body)
5933
}
6034

6135
// Copied from rustc_mir_transform/src/lib.rs.
@@ -148,7 +122,7 @@ fn take_array<T, const N: usize>(b: &mut Box<[T]>) -> Result<[T; N], Box<[T]>> {
148122
memoize!(
149123
pub fn analysis_mir<'tcx>(cx: &AnalysisCtxt<'tcx>, def_id: DefId) -> &'tcx Body<'tcx> {
150124
if let Some(local_def_id) = def_id.as_local() {
151-
local_analysis_mir(cx.tcx, local_def_id)
125+
local_analysis_mir(cx, local_def_id)
152126
} else if let Some(mir) = cx.sql_load_with_span::<analysis_mir>(def_id, cx.def_span(def_id))
153127
{
154128
mir

0 commit comments

Comments
 (0)