1use super::{cache::{DiskCache, Record, now_seconds}, policy::*};
5use std::{cell::Cell, collections::{BTreeMap, BTreeSet, VecDeque}, path::PathBuf,
6 string::{String, ToString}, sync::{Mutex, MutexGuard}, time::{Duration, Instant}, vec::Vec};
7
8std::thread_local! { static DEPTH: Cell<usize> = const { Cell::new(0) }; }
9pub fn is_tuning() -> bool { DEPTH.with(|d| d.get() != 0) }
10pub(super) struct DepthGuard;
11impl DepthGuard { pub(super) fn enter() -> Self { DEPTH.with(|d| d.set(d.get()+1)); Self } }
12impl Drop for DepthGuard { fn drop(&mut self) { DEPTH.with(|d| d.set(d.get()-1)); } }
13
14pub trait TrialRunner {
17 fn validate(&mut self, reference: usize, candidate: usize, tolerance: Tolerance) -> Result<Validation, TuneFailure>;
18 fn measure(&mut self, candidate: usize) -> Result<Duration, TuneFailure>;
19}
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub enum DecisionSource { Tuned, MemoryCache, DiskCache, CacheMiss, Disabled, Busy, Nested, ValidationUnavailable }
22#[derive(Debug, Clone)]
23pub struct Decision {
24 pub index: usize,
25 pub reference_index: usize,
26 pub name: String,
27 pub source: DecisionSource,
28 pub verified: bool,
30 pub ratio: Option<f64>,
31 pub cache_key: String,
32}
33#[derive(Debug, Clone)]
34pub struct CandidateReport {
35 pub name: String,
36 pub samples: usize,
37 pub ratio: Option<f64>,
38 pub relative_mad: Option<f64>,
39 pub verified: bool,
40 pub note: String,
41}
42#[derive(Debug, Clone)]
43pub struct TuneReport {
44 pub operation: String,
45 pub environment: String,
46 pub workload: String,
47 pub execution_context: String,
48 pub winner: String,
49 pub elapsed: Duration,
50 pub budget_exhausted: bool,
51 pub candidates: Vec<CandidateReport>,
52}
53#[derive(Debug, Clone, Default)]
54pub struct Stats {
55 pub memory_hits: u64, pub disk_hits: u64, pub misses: u64,
56 pub busy_fallbacks: u64, pub tunes: u64, pub cache_warnings: u64,
57 pub invalidations: u64, pub device_failures: u64,
58}
59#[derive(Default)]
60struct State {
61 records: BTreeMap<String, Record>, order: VecDeque<String>,
62 bypass: BTreeMap<String, Scope>,
63 pending: BTreeSet<String>, devices: BTreeSet<String>, poisoned_devices: BTreeSet<String>,
64 banned: BTreeMap<String, BTreeSet<String>>, regressions: BTreeMap<String, VecDeque<f64>>,
65 reports: VecDeque<TuneReport>, stats: Stats,
66}
67pub struct StackTuner { policy: StackPolicy, disk: Option<DiskCache>, state: Mutex<State> }
70struct Permit<'a> { tuner: &'a StackTuner, key: String, device: String }
71impl Drop for Permit<'_> {
72 fn drop(&mut self) { let mut s = self.tuner.lock(); s.pending.remove(&self.key); s.devices.remove(&self.device); }
73}
74impl StackTuner {
75 pub fn new(policy: StackPolicy, cache_directory: Option<PathBuf>) -> Result<Self, TuneFailure> {
76 policy.validate()?;
77 let disk = cache_directory.map(|dir| DiskCache::new(dir, policy.capacity));
78 Ok(Self { policy, disk, state: Mutex::new(State::default()) })
79 }
80 fn lock(&self) -> MutexGuard<'_, State> { self.state.lock().unwrap_or_else(|e| e.into_inner()) }
81 pub fn policy(&self) -> &StackPolicy { &self.policy }
82 pub fn stats(&self) -> Stats { self.lock().stats.clone() }
83 pub fn lower_level_fingerprint(&self) -> String {
87 let s = self.lock();
88 let mut parts = Vec::new();
89 for (key, record) in &s.records {
90 if record.scope != Scope::Pipeline as u8 && record.is_fresh(now_seconds(), self.policy.ttl.as_secs().max(1)) { parts.push(fields(&[key, &record.winner])); }
91 }
92 for (key, scope) in &s.bypass {
94 if *scope != Scope::Pipeline { parts.push(fields(&[key, "unverified-reference"])); }
95 }
96 parts.sort();
97 super::cache::digest(fields(&parts.iter().map(String::as_str).collect::<Vec<_>>()).as_bytes())
98 }
99 pub fn reports(&self) -> Vec<TuneReport> { self.lock().reports.iter().cloned().collect() }
100 fn insert(&self, record: Record) {
101 let mut s = self.lock();
102 if !s.order.contains(&record.key) {
103 while s.order.len() >= self.policy.capacity {
104 if let Some(key) = s.order.pop_front() {
105 s.records.remove(&key); s.banned.remove(&key); s.regressions.remove(&key); s.bypass.remove(&key);
106 }
107 }
108 s.order.push_back(record.key.clone());
109 }
110 s.records.insert(record.key.clone(), record);
111 }
112 fn allowed(&self, key: &str, c: &Candidate) -> bool {
113 c.fits(&self.policy) && !self.lock().banned.get(key).is_some_and(|b| b.contains(&c.name))
114 }
115 fn from_record(&self, r: &Record, candidates: &[Candidate], reference: usize, source: DecisionSource) -> Option<Decision> {
116 if !r.is_fresh(now_seconds(), self.policy.ttl.as_secs().max(1)) || (self.policy.require_validation && !r.verified) { return None; }
117 let index = candidates.iter().position(|c| c.name == r.winner && self.allowed(&r.key, c))?;
118 Some(Decision { index, reference_index: reference, name: r.winner.clone(), source, verified: r.verified, ratio: Some(r.ratio), cache_key: r.key.clone() })
119 }
120 fn validation(&self, runner: &mut impl TrialRunner, reference: usize, candidate: usize) -> Result<bool, TuneFailure> {
121 match runner.validate(reference, candidate, self.policy.tolerance)? {
122 Validation::Passed => Ok(true),
123 Validation::Unsupported if !self.policy.require_validation => Ok(false),
124 Validation::Unsupported => Err(TuneFailure { kind: FailureKind::Unavailable, message: "no numerical validator for this output/size".into() }),
125 }
126 }
127 fn fail_device(&self, environment: &str) {
128 let mut s = self.lock(); s.poisoned_devices.insert(environment.to_string()); s.stats.device_failures += 1;
129 }
130 pub fn select(&self, problem: &Problem, candidates: &[Candidate], reference: usize, runner: &mut impl TrialRunner) -> Result<Decision, TuneFailure> {
131 if candidates.is_empty() || candidates.len() > 4096 || reference >= candidates.len()
132 || problem.operation.is_empty() || problem.environment.is_empty() || problem.workload.is_empty() {
133 return Err(TuneFailure::invalid("a workload, environment and valid reference candidate are required"));
134 }
135 let mut names = BTreeSet::new();
136 if candidates.iter().any(|c| c.name.is_empty() || c.name.len() > 4096 || !names.insert(&c.name)) {
137 return Err(TuneFailure::invalid("candidate names must be nonempty, bounded and unique"));
138 }
139 if !candidates[reference].fits(&self.policy) { return Err(TuneFailure::invalid("reference does not satisfy eligibility/workspace policy")); }
140 let key = cache_key(problem, candidates, reference, &self.policy);
141 if key.len() > 384 * 1024 { return Err(TuneFailure::invalid("autotune key exceeds 384 KiB")); }
142 let fallback = |source| Decision { index: reference, reference_index: reference, name: candidates[reference].name.clone(), source, verified: false, ratio: None, cache_key: key.clone() };
143 if self.lock().poisoned_devices.contains(&problem.environment) {
144 return Err(TuneFailure { kind: FailureKind::Quarantined, message: "device tuning lane is quarantined after an unconfirmed device completion".into() });
145 }
146 if self.policy.mode == Mode::Disabled { return Ok(fallback(DecisionSource::Disabled)); }
147 let memory = { self.lock().records.get(&key).cloned() };
148 if let Some(record) = memory {
149 if let Some(d) = self.from_record(&record, candidates, reference, DecisionSource::MemoryCache) {
150 self.lock().stats.memory_hits += 1; return Ok(d);
151 }
152 }
153 if self.lock().bypass.contains_key(&key) { return Ok(fallback(DecisionSource::ValidationUnavailable)); }
154 if is_tuning() { return Ok(fallback(DecisionSource::Nested)); }
155 let _permit = {
156 let mut s = self.lock();
157 if s.pending.contains(&key) || s.devices.contains(&problem.environment) || s.pending.len() >= self.policy.max_parallel_tunes {
158 s.stats.busy_fallbacks += 1; return Ok(fallback(DecisionSource::Busy));
159 }
160 s.pending.insert(key.clone()); s.devices.insert(problem.environment.clone()); s.stats.misses += 1;
161 Permit { tuner: self, key: key.clone(), device: problem.environment.clone() }
162 };
163 let _depth = DepthGuard::enter();
164 if problem.persistent {
165 if let Some(disk) = &self.disk {
166 match disk.load(&key) {
167 Ok(Some(record)) => {
168 if let Some(mut decision) = self.from_record(&record, candidates, reference, DecisionSource::DiskCache).filter(|_| record.scope == problem.scope as u8) {
169 match self.validation(runner, reference, decision.index) {
172 Ok(checked) => {
173 decision.verified = checked;
174 let mut record = record; record.verified = checked;
175 self.insert(record); self.lock().stats.disk_hits += 1; return Ok(decision);
176 }
177 Err(e) if e.kind == FailureKind::Device => { self.fail_device(&problem.environment); return Err(e); }
178 Err(_) => { if disk.remove(&key).is_err() { self.lock().stats.cache_warnings += 1; } }
179 }
180 }
181 }
182 Err(_) => self.lock().stats.cache_warnings += 1,
183 Ok(None) => {}
184 }
185 }
186 }
187 if self.policy.mode == Mode::CacheOnly { return Ok(fallback(DecisionSource::CacheMiss)); }
188 match self.explore(problem, candidates, reference, &key, runner) {
189 Ok((decision, report, record)) => {
190 if problem.persistent {
191 if let Some(disk) = &self.disk { if disk.save(&record).is_err() { self.lock().stats.cache_warnings += 1; } }
192 }
193 self.insert(record);
194 let mut s = self.lock(); s.stats.tunes += 1;
195 if s.reports.len() >= self.policy.capacity.min(128) { s.reports.pop_front(); }
196 s.reports.push_back(report);
197 Ok(decision)
198 }
199 Err(e) if e.kind == FailureKind::Unavailable => {
200 let mut s = self.lock();
203 if !s.order.contains(&key) { s.order.push_back(key.clone()); }
204 s.bypass.insert(key.clone(), problem.scope);
205 while s.order.len() > self.policy.capacity {
206 if let Some(old) = s.order.pop_front() { s.records.remove(&old); s.banned.remove(&old); s.regressions.remove(&old); s.bypass.remove(&old); }
207 }
208 if s.reports.len() >= self.policy.capacity.min(128) { s.reports.pop_front(); }
209 s.reports.push_back(TuneReport {
210 operation: problem.operation.clone(), environment: problem.environment.clone(),
211 workload: problem.workload.clone(), execution_context: problem.execution_context.clone(),
212 winner: candidates[reference].name.clone(), elapsed: Duration::ZERO, budget_exhausted: false,
213 candidates: std::vec![CandidateReport { name: candidates[reference].name.clone(), samples: 0,
214 ratio: None, relative_mad: None, verified: false, note: e.message }],
215 });
216 Ok(fallback(DecisionSource::ValidationUnavailable))
217 }
218 Err(e) => { if e.kind == FailureKind::Device { self.fail_device(&problem.environment); } Err(e) }
219 }
220 }
221 fn explore(&self, problem: &Problem, candidates: &[Candidate], reference: usize, key: &str, runner: &mut impl TrialRunner) -> Result<(Decision, TuneReport, Record), TuneFailure> {
222 let started = Instant::now();
223 for _ in 0..self.policy.warmups { runner.measure(reference)?; }
226 let reference_checked = self.validation(runner, reference, reference)?;
227 let reference_time = runner.measure(reference)?;
228 if reference_time.is_zero() { return Err(TuneFailure { kind: FailureKind::Unavailable, message: "reference has zero elapsed time".into() }); }
229 let mut winner = reference; let mut winner_ratio = 1.0; let mut winner_time = reference_time;
230 let mut winner_reference = reference_time; let mut winner_checked = reference_checked;
231 let mut reports = Vec::new();
232 reports.push(CandidateReport { name: candidates[reference].name.clone(), samples: 1, ratio: Some(1.0), relative_mad: Some(0.0), verified: reference_checked, note: "reference; retained unless a stable measured improvement qualifies".into() });
233 let mut attempted = 1; let mut exhausted = false;
234 for (index, candidate) in candidates.iter().enumerate() {
235 if index == reference { continue; }
236 let mut report = CandidateReport { name: candidate.name.clone(), samples: 0, ratio: None, relative_mad: None, verified: false, note: String::new() };
237 if !self.allowed(key, candidate) { report.note = "ineligible, unknown/over-budget workspace, or regressed candidate".into(); reports.push(report); continue; }
238 if attempted >= self.policy.max_candidates || started.elapsed() >= self.policy.budget {
239 exhausted = true; report.note = "search budget exhausted".into(); reports.push(report); continue;
240 }
241 attempted += 1;
242 let result = (|| -> Result<(bool, Vec<(Duration, Duration)>), TuneFailure> {
243 for _ in 0..self.policy.warmups {
244 if started.elapsed() >= self.policy.budget { return Err(TuneFailure::rejected("budget exhausted during warmup")); }
245 runner.measure(index)?;
246 }
247 let checked = self.validation(runner, reference, index)?;
248 let mut pairs = Vec::with_capacity(self.policy.samples);
249 for sample in 0..self.policy.samples {
250 if started.elapsed() >= self.policy.budget { return Err(TuneFailure::rejected("budget exhausted before enough paired samples")); }
251 let pair = if (sample + index) % 2 == 0 {
253 let a = runner.measure(reference)?; let b = runner.measure(index)?; (a,b)
254 } else {
255 let b = runner.measure(index)?; let a = runner.measure(reference)?; (a,b)
256 };
257 if pair.0.is_zero() || pair.1.is_zero() { return Err(TuneFailure::rejected("zero duration cannot be scored")); }
258 pairs.push(pair);
259 }
260 Ok((checked, pairs))
261 })();
262 match result {
263 Err(e) if e.kind == FailureKind::Device => return Err(e),
264 Err(e) => { report.note = e.message; },
265 Ok((checked, pairs)) => {
266 report.verified = checked; report.samples = pairs.len();
267 match paired_score(&pairs, &self.policy) {
268 None => { report.note = "timings too noisy or incomplete".into(); },
269 Some((ratio, mad)) => {
270 report.ratio = Some(ratio); report.relative_mad = Some(mad);
271 report.note = "paired synchronized measurements".into();
272 if ratio < winner_ratio && ratio * self.policy.min_speedup <= 1.0 {
273 winner = index; winner_ratio = ratio; winner_checked = checked;
274 let mut a: Vec<_> = pairs.iter().map(|p| p.0.as_secs_f64()).collect();
275 let mut b: Vec<_> = pairs.iter().map(|p| p.1.as_secs_f64()).collect();
276 winner_reference = Duration::from_secs_f64(median(&mut a).unwrap());
277 winner_time = Duration::from_secs_f64(median(&mut b).unwrap());
278 }
279 }
280 }
281 }
282 }
283 reports.push(report);
284 }
285 exhausted |= started.elapsed() >= self.policy.budget;
286 let record = Record { scope: problem.scope as u8, key: key.to_string(), winner: candidates[winner].name.clone(), created: now_seconds(),
287 reference_ns: winner_reference.as_nanos().min(u64::MAX as u128) as u64,
288 winner_ns: winner_time.as_nanos().min(u64::MAX as u128) as u64, ratio: winner_ratio, verified: winner_checked };
289 let decision = Decision { index: winner, reference_index: reference, name: record.winner.clone(), source: DecisionSource::Tuned, verified: winner_checked, ratio: Some(winner_ratio), cache_key: key.to_string() };
290 let report = TuneReport { operation: problem.operation.clone(), environment: problem.environment.clone(), workload: problem.workload.clone(), execution_context: problem.execution_context.clone(), winner: record.winner.clone(), elapsed: started.elapsed(), budget_exhausted: exhausted, candidates: reports };
291 Ok((decision, report, record))
292 }
293 pub fn invalidate(&self, decision: &Decision, ban_candidate: bool) {
295 {
296 let mut s = self.lock(); s.records.remove(&decision.cache_key); s.regressions.remove(&decision.cache_key); s.bypass.remove(&decision.cache_key);
297 if ban_candidate && decision.index != decision.reference_index {
298 s.banned.entry(decision.cache_key.clone()).or_default().insert(decision.name.clone());
299 }
300 if !s.order.contains(&decision.cache_key) { s.order.push_back(decision.cache_key.clone()); }
302 while s.order.len() > self.policy.capacity {
303 if let Some(old) = s.order.pop_front() { s.records.remove(&old); s.banned.remove(&old); s.regressions.remove(&old); s.bypass.remove(&old); }
304 }
305 s.stats.invalidations += 1;
306 }
307 if let Some(disk) = &self.disk { if disk.remove(&decision.cache_key).is_err() { self.lock().stats.cache_warnings += 1; } }
308 }
309 pub fn record_comparison(&self, decision: &Decision, reference: Duration, selected: Duration, correctness_checked: bool) -> Result<bool, TuneFailure> {
312 if !correctness_checked || reference.is_zero() || selected.is_zero() { return Err(TuneFailure::invalid("regression observations require nonzero, correctness-checked paired timings")); }
313 let should_invalidate = {
314 let mut s = self.lock();
315 if !s.records.get(&decision.cache_key).is_some_and(|r| r.winner == decision.name) { return Ok(false); }
316 let history = s.regressions.entry(decision.cache_key.clone()).or_default();
317 if history.len() == self.policy.regression_pairs { history.pop_front(); }
318 history.push_back(selected.as_secs_f64()/reference.as_secs_f64());
319 if history.len() < self.policy.regression_pairs { false } else {
320 let mut ratios: Vec<_> = history.iter().copied().collect(); median(&mut ratios).is_some_and(|r| r > self.policy.regression_ratio)
321 }
322 };
323 if should_invalidate { self.invalidate(decision, true); }
324 Ok(should_invalidate)
325 }
326}