Skip to main content

flux_refineck/
lib.rs

1//! Refinement type checking
2
3#![feature(
4    associated_type_defaults,
5    box_patterns,
6    min_specialization,
7    never_type,
8    rustc_private,
9    unwrap_infallible
10)]
11
12extern crate rustc_abi;
13extern crate rustc_data_structures;
14extern crate rustc_errors;
15extern crate rustc_hir;
16extern crate rustc_index;
17extern crate rustc_infer;
18extern crate rustc_middle;
19extern crate rustc_mir_dataflow;
20extern crate rustc_span;
21extern crate rustc_type_ir;
22
23mod checker;
24pub mod compare_impl_item;
25mod ghost_statements;
26pub mod invariants;
27mod primops;
28mod queue;
29mod type_env;
30
31use checker::{Checker, trait_impl_subtyping};
32use flux_common::{dbg, dbg::SpanTrace, result::ResultExt as _};
33use flux_config as config;
34use flux_infer::{
35    fixpoint_encoding::{
36        FixQueryCache, FixpointCheckError, PossibleSolutions, SolutionTrace, TagIdx,
37    },
38    infer::{ConstrReason, SubtypeReason, Tag},
39    wkvars::WKVarSubst,
40};
41use flux_macros::msg;
42use flux_middle::{
43    FixpointQueryKind,
44    def_id::MaybeExternId,
45    global_env::GlobalEnv,
46    metrics::{self, Metric, TimingKind},
47    pretty,
48    rty::{self, ESpan, EarlyBinder, fold::TypeFoldable},
49};
50use rustc_data_structures::{fx::FxHashMap, unord::UnordMap};
51use rustc_errors::{Applicability, Diag, ErrorGuaranteed};
52use rustc_hir::def_id::{DefId, LocalDefId};
53use rustc_span::Span;
54
55use crate::{checker::errors::ResultExt as _, ghost_statements::compute_ghost_statements};
56
57pub fn report_fixpoint_errors(
58    genv: GlobalEnv,
59    local_id: LocalDefId,
60    errors: Vec<FixpointCheckError<Tag>>,
61) -> Result<(), ErrorGuaranteed> {
62    #[expect(clippy::collapsible_else_if, reason = "it looks better")]
63    if genv.should_fail(local_id) {
64        if errors.is_empty() { report_expected_neg(genv, local_id) } else { Ok(()) }
65    } else {
66        if errors.is_empty() { Ok(()) } else { report_errors(genv, local_id, errors) }
67    }
68}
69
70fn check_body(
71    genv: GlobalEnv,
72    cache: &mut FixQueryCache,
73    def_id: LocalDefId,
74    poly_sig: &rty::PolyFnSig,
75) -> Result<(), ErrorGuaranteed> {
76    let span = genv.tcx().def_span(def_id);
77    let opts = genv.infer_opts(def_id);
78
79    dbg::log_verbose!("FLUX checking: {def_id:?} {span:?}");
80
81    let ghost_stmts = compute_ghost_statements(genv, def_id)
82        .with_span(span)
83        .map_err(|err| err.emit(genv, def_id))?;
84    let mut closures = UnordMap::default();
85
86    // PHASE 1: infer shape of `TypeEnv` at the entry of join points
87    let shape_result =
88        Checker::run_in_shape_mode(genv, def_id, &ghost_stmts, &mut closures, opts, poly_sig)
89            .map_err(|err| err.emit(genv, def_id))?;
90
91    // PHASE 2: generate refinement tree constraint
92    let infcx_root = Checker::run_in_refine_mode(
93        genv,
94        def_id,
95        &ghost_stmts,
96        &mut closures,
97        shape_result,
98        opts,
99        poly_sig,
100    )
101    .map_err(|err| err.emit(genv, def_id))?;
102
103    // PHASE 3: invoke fixpoint on the constraint
104    if (genv.proven_externally(def_id).is_some() && flux_config::lean().is_check())
105        || flux_config::lean().is_emit()
106    {
107        infcx_root
108            .execute_lean_query(cache, MaybeExternId::Local(def_id))
109            .emit(&genv)
110    } else {
111        let answer = infcx_root
112            .execute_fixpoint_query(cache, MaybeExternId::Local(def_id), FixpointQueryKind::Body)
113            .emit(&genv)?;
114
115        let tcx = genv.tcx();
116        let hir_id = tcx.local_def_id_to_hir_id(def_id);
117        let body_span = tcx.hir_span_with_body(hir_id);
118        dbg::solution!(genv, &answer, body_span);
119
120        let errors = answer.errors;
121        report_fixpoint_errors(genv, def_id, errors)
122    }
123}
124
125pub fn check_static(
126    genv: GlobalEnv,
127    cache: &mut FixQueryCache,
128    def_id: LocalDefId,
129    ty: rty::Ty,
130) -> Result<(), ErrorGuaranteed> {
131    // Build a PolyFnSig with no inputs and `ty` as the output
132    let output = rty::Binder::dummy(rty::FnOutput::new(ty, vec![]));
133    let fn_sig = rty::FnSig::new(
134        rustc_hir::Safety::Safe,
135        rustc_abi::ExternAbi::Rust,
136        rty::List::empty(),
137        rty::List::empty(),
138        output,
139        rty::Expr::ff(),
140        false,
141    );
142    let poly_sig = rty::PolyFnSig::dummy(fn_sig);
143
144    metrics::incr_metric(Metric::FnChecked, 1);
145    metrics::time_it(TimingKind::CheckBody(def_id), || check_body(genv, cache, def_id, &poly_sig))
146}
147
148pub fn check_fn(
149    genv: GlobalEnv,
150    cache: &mut FixQueryCache,
151    def_id: LocalDefId,
152) -> Result<(), ErrorGuaranteed> {
153    let span = genv.tcx().def_span(def_id);
154
155    // Code generated by a `#[derive(..)]` can't be annotated with `#[trusted]` directly, so a
156    // type opts its derived code out of checking with `#[flux::trusted_derive]`. This is needed
157    // for types whose derives flux cannot handle, e.g. a `#[flux::opaque]` struct whose derived
158    // `Debug`/`Hash` read the internal representation.
159    if span.in_derive_expansion()
160        && let Some(adt_def_id) = genv.derive_self_ty(def_id)
161        && genv.trusted_derive(adt_def_id)
162    {
163        metrics::incr_metric(Metric::FnTrusted, 1);
164        return Ok(());
165    }
166
167    let opts = genv.infer_opts(def_id);
168
169    // FIXME(nilehmann) we should move this check to `compare_impl_item`
170    if let Some(infcx_root) = trait_impl_subtyping(genv, def_id, opts, span)
171        .with_span(span)
172        .map_err(|err| err.emit(genv, def_id))?
173    {
174        tracing::info!("check_fn::refine-subtyping");
175        let answer = infcx_root
176            .execute_fixpoint_query(cache, MaybeExternId::Local(def_id), FixpointQueryKind::Impl)
177            .emit(&genv)?;
178        tracing::info!("check_fn::fixpoint-subtyping");
179        let errors = answer.errors;
180        report_fixpoint_errors(genv, def_id, errors)?;
181    }
182
183    // Skip trusted functions
184    if genv.trusted(def_id) {
185        metrics::incr_metric(Metric::FnTrusted, 1);
186        return Ok(());
187    }
188
189    metrics::incr_metric(Metric::FnChecked, 1);
190    metrics::time_it(TimingKind::CheckBody(def_id), || -> Result<(), ErrorGuaranteed> {
191        let poly_sig = genv
192            .fn_sig(def_id)
193            .with_span(span)
194            .map_err(|err| err.emit(genv, def_id))?
195            .instantiate_identity();
196        let poly_sig = rty::auto_strong(genv, def_id, poly_sig);
197
198        check_body(genv, cache, def_id, &poly_sig)
199    })?;
200
201    dbg::check_fn_span!(genv.tcx(), def_id).in_scope(|| Ok(()))
202}
203
204fn call_error<'a>(genv: GlobalEnv<'a, '_>, span: Span, dst_span: Option<ESpan>) -> Diag<'a> {
205    genv.sess()
206        .dcx()
207        .handle()
208        .create_err(errors::RefineError::call(span, dst_span))
209}
210
211fn ret_error<'a>(genv: GlobalEnv<'a, '_>, span: Span, dst_span: Option<ESpan>) -> Diag<'a> {
212    genv.sess()
213        .dcx()
214        .handle()
215        .create_err(errors::RefineError::ret(span, dst_span))
216}
217
218fn report_errors(
219    genv: GlobalEnv,
220    local_id: LocalDefId,
221    errors: Vec<FixpointCheckError<Tag>>,
222) -> Result<(), ErrorGuaranteed> {
223    let log_path = if config::dump_constraint() {
224        let path = dbg::item_dump_path(genv.tcx(), local_id.to_def_id(), "smt2");
225        if path.exists() { Some(path) } else { None }
226    } else {
227        None
228    };
229    let mut solutions_by_tag: FxHashMap<Tag, (TagIdx, PossibleSolutions)> = FxHashMap::default();
230    for error in errors {
231        if let Some((_, val)) = solutions_by_tag.get_mut(&error.tag) {
232            val.extend(error.possible_solutions);
233        } else {
234            solutions_by_tag.insert(error.tag, (error.tag_idx, error.possible_solutions));
235        }
236    }
237    let rerun_note = rerun_hint_note(genv, local_id);
238    let mut e = None;
239    for (tag, (tag_idx, possible_solutions)) in solutions_by_tag {
240        let span = tag.src_span;
241        let mut err_diag = match tag.reason {
242            ConstrReason::Call
243            | ConstrReason::Subtype(SubtypeReason::Input)
244            | ConstrReason::Subtype(SubtypeReason::Requires)
245            | ConstrReason::Predicate => call_error(genv, span, tag.dst_span),
246            ConstrReason::Assign => {
247                genv.sess()
248                    .dcx()
249                    .handle()
250                    .create_err(errors::AssignError { span })
251            }
252            ConstrReason::Ret
253            | ConstrReason::Subtype(SubtypeReason::Output)
254            | ConstrReason::Subtype(SubtypeReason::Ensures) => ret_error(genv, span, tag.dst_span),
255            ConstrReason::Div => {
256                genv.sess()
257                    .dcx()
258                    .handle()
259                    .create_err(errors::DivError { span })
260            }
261            ConstrReason::Rem => {
262                genv.sess()
263                    .dcx()
264                    .handle()
265                    .create_err(errors::RemError { span })
266            }
267            ConstrReason::Goto(_) => {
268                genv.sess()
269                    .dcx()
270                    .handle()
271                    .create_err(errors::GotoError { span })
272            }
273            ConstrReason::Assert(msg) => {
274                genv.sess()
275                    .dcx()
276                    .handle()
277                    .create_err(errors::AssertError { span, msg })
278            }
279            ConstrReason::Fold | ConstrReason::FoldLocal => {
280                genv.sess()
281                    .dcx()
282                    .handle()
283                    .create_err(errors::FoldError::new(span, tag.dst_span))
284            }
285            ConstrReason::Overflow => {
286                genv.sess()
287                    .dcx()
288                    .handle()
289                    .create_err(errors::OverflowError { span })
290            }
291            ConstrReason::Underflow => {
292                genv.sess()
293                    .dcx()
294                    .handle()
295                    .create_err(errors::UnderflowError { span })
296            }
297            ConstrReason::Other => {
298                genv.sess()
299                    .dcx()
300                    .handle()
301                    .create_err(errors::UnknownError { span })
302            }
303            ConstrReason::NoPanic(callee, reason) => {
304                genv.sess().dcx().handle().create_err(errors::PanicError {
305                    span,
306                    callee: genv.tcx().def_path_debug_str(callee),
307                    reason: format!("{:?}", reason),
308                })
309            }
310        };
311        let wkvar_solutions = possible_solutions
312            .iter()
313            .flat_map(|(wkvid, solutions)| solutions.iter().map(move |solution| (wkvid, solution)));
314        for (wkvid, solution) in wkvar_solutions {
315            add_fn_fix_diagnostic(genv, &mut err_diag, wkvid.clone(), solution);
316        }
317        if let Some(note) = &rerun_note {
318            err_diag.note(note.clone());
319        }
320        if let Some(path) = &log_path {
321            err_diag.arg("path", path.display().to_string());
322            err_diag.arg("tag", tag_idx.to_string());
323            err_diag.note(msg!("log file saved to {$path} (tag: {$tag})"));
324        }
325        e = Some(err_diag.emit());
326    }
327
328    if let Some(e) = e { Err(e) } else { Ok(()) }
329}
330
331fn report_expected_neg(genv: GlobalEnv, def_id: LocalDefId) -> Result<(), ErrorGuaranteed> {
332    Err(genv.sess().emit_err(errors::ExpectedNeg {
333        span: genv.tcx().def_span(def_id),
334        def_descr: genv.tcx().def_descr(def_id.to_def_id()),
335    }))
336}
337
338fn add_fn_fix_diagnostic<'a>(
339    genv: GlobalEnv<'a, '_>,
340    diag: &mut Diag<'a>,
341    wkvid: rty::WKVid,
342    solution: &rty::Binder<rty::Expr>,
343) {
344    let pretty_solution = solution.map_ref(|e| e.simplify(&Default::default()).prettify());
345    let fn_sig = genv.fn_sig(wkvid.parent_fn).unwrap();
346    let mut wkvar_subst = WKVarSubst::new(
347        std::iter::once((wkvid.clone(), pretty_solution)).collect::<UnordMap<_, _>>(),
348        false,
349    );
350    let solved_fn_sig = EarlyBinder(fn_sig.skip_binder_ref().fold_with(&mut wkvar_subst));
351    let fixed_fn_sig_snippet = format!(
352        "{:?}",
353        pretty::with_cx!(&pretty::PrettyCx::default(genv).hide_regions(true), &solved_fn_sig)
354    );
355    let fn_first_line = fn_first_line(genv, wkvid.parent_fn);
356    let fn_first_line_snippet = genv
357        .tcx()
358        .sess
359        .source_map()
360        .span_to_snippet(fn_first_line)
361        .unwrap_or_else(|_| panic!("No snippet for span {:?}", fn_first_line));
362    let prefix_spaces = &fn_first_line_snippet[..fn_first_line_snippet
363        .find(|c: char| !c.is_whitespace())
364        .unwrap_or(fn_first_line_snippet.len())];
365    let subst_solutions = &wkvar_subst.subst_instantiations[&wkvid];
366    assert!(subst_solutions.len() == 1);
367
368    // Check if there's an existing spec attribute that needs to be replaced
369    if let Some(old_spec_span) = genv.spec_attr_span(wkvid.parent_fn) {
370        diag.span_suggestion(
371            old_spec_span,
372            "try replacing the refinement",
373            format!("{}#[flux_rs::sig({})]", prefix_spaces, fixed_fn_sig_snippet),
374            Applicability::MachineApplicable,
375        );
376    } else {
377        diag.span_suggestion(
378            fn_first_line,
379            "try adding the refinement",
380            format!(
381                "{}#[flux_rs::sig({})]\n{}",
382                prefix_spaces, fixed_fn_sig_snippet, fn_first_line_snippet
383            ),
384            Applicability::MachineApplicable,
385        );
386    }
387}
388
389fn fn_first_line<'a>(genv: GlobalEnv<'a, '_>, def_id: DefId) -> Span {
390    let span = genv.tcx().def_span(def_id);
391    let first_line = genv
392        .tcx()
393        .sess
394        .source_map()
395        .lookup_line(span.lo())
396        .unwrap_or_else(|_| panic!("span for {:?} doesn't have a first line", def_id));
397    let first_line_range = first_line.sf.line_bounds(first_line.line);
398    Span::new(first_line_range.start, first_line_range.end, span.ctxt(), None)
399}
400
401fn rerun_hint_note(genv: GlobalEnv, def_id: LocalDefId) -> Option<String> {
402    if !config::rerun_hint() || !config::inside_cargo() {
403        return None;
404    }
405    let pattern = format!("def:{}", genv.tcx().def_path_str(def_id));
406    let pkg = std::env::var("CARGO_PKG_NAME")
407        .map(|p| format!(" -p {p}"))
408        .unwrap_or_default();
409    Some(format!("to rerun: `cargo flux check{pkg} --only-check={}`", shell_quote_arg(&pattern)))
410}
411
412fn shell_quote_arg(arg: &str) -> String {
413    format!("'{}'", arg.replace('\'', "'\\''"))
414}
415
416mod errors {
417    use flux_errors::E0999;
418    use flux_macros::{Diagnostic, Subdiagnostic};
419    use flux_middle::rty::ESpan;
420    use rustc_span::Span;
421
422    #[derive(Diagnostic)]
423    #[diag("error jumping to join point", code = E0999)]
424    pub struct GotoError {
425        #[primary_span]
426        pub span: Span,
427    }
428
429    #[derive(Diagnostic)]
430    #[diag("assignment might be unsafe", code = E0999)]
431    pub struct AssignError {
432        #[primary_span]
433        pub span: Span,
434    }
435
436    #[derive(Subdiagnostic)]
437    #[note("this is the condition that cannot be proved")]
438    pub(crate) struct ConditionSpanNote {
439        #[primary_span]
440        pub span: Span,
441    }
442
443    #[derive(Subdiagnostic)]
444    #[note("inside this call")]
445    pub(crate) struct CallSpanNote {
446        #[primary_span]
447        pub span: Span,
448    }
449
450    #[derive(Diagnostic)]
451    #[diag("refinement type error", code = E0999)]
452    pub struct RefineError {
453        #[primary_span]
454        #[label("a {$cond} cannot be proved")]
455        pub span: Span,
456        cond: &'static str,
457        #[subdiagnostic]
458        span_note: Option<ConditionSpanNote>,
459        #[subdiagnostic]
460        call_span_note: Option<CallSpanNote>,
461    }
462
463    impl RefineError {
464        pub fn call(span: Span, espan: Option<ESpan>) -> Self {
465            RefineError::new("precondition", span, espan)
466        }
467
468        pub fn ret(span: Span, espan: Option<ESpan>) -> Self {
469            RefineError::new("postcondition", span, espan)
470        }
471
472        fn new(cond: &'static str, span: Span, espan: Option<ESpan>) -> RefineError {
473            match espan {
474                Some(dst_span) => {
475                    let span_note = Some(ConditionSpanNote { span: dst_span.span });
476                    let call_span_note = dst_span.base.map(|span| CallSpanNote { span });
477                    RefineError { span, cond, span_note, call_span_note }
478                }
479                None => RefineError { span, cond, span_note: None, call_span_note: None },
480            }
481        }
482    }
483
484    #[derive(Diagnostic)]
485    #[diag("possible division by zero", code = E0999)]
486    pub struct DivError {
487        #[primary_span]
488        pub span: Span,
489    }
490
491    #[derive(Diagnostic)]
492    #[diag("possible remainder with a divisor of zero", code = E0999)]
493    pub struct RemError {
494        #[primary_span]
495        pub span: Span,
496    }
497
498    #[derive(Diagnostic)]
499    #[diag("assertion might fail: {$msg}", code = E0999)]
500    pub struct AssertError {
501        #[primary_span]
502        pub span: Span,
503        pub msg: &'static str,
504    }
505
506    #[derive(Diagnostic)]
507    #[diag("type invariant may not hold (when place is folded)", code = E0999)]
508    pub struct FoldError {
509        #[primary_span]
510        pub span: Span,
511        #[subdiagnostic]
512        span_note: Option<ConditionSpanNote>,
513    }
514
515    impl FoldError {
516        pub fn new(span: Span, espan: Option<ESpan>) -> Self {
517            let span_note = espan.map(|espan| ConditionSpanNote { span: espan.span });
518            FoldError { span, span_note }
519        }
520    }
521
522    #[derive(Diagnostic)]
523    #[diag("arithmetic operation may overflow", code = E0999)]
524    pub struct OverflowError {
525        #[primary_span]
526        pub span: Span,
527    }
528
529    #[derive(Diagnostic)]
530    #[diag("arithmetic operation may underflow", code = E0999)]
531    pub struct UnderflowError {
532        #[primary_span]
533        pub span: Span,
534    }
535
536    #[derive(Diagnostic)]
537    #[diag("cannot prove this code safe", code = E0999)]
538    pub struct UnknownError {
539        #[primary_span]
540        pub span: Span,
541    }
542
543    #[derive(Diagnostic)]
544    #[diag("{$def_descr} marked with `#[should_fail]` didn't produce a refinement type error", code = E0999)]
545    pub struct ExpectedNeg {
546        #[primary_span]
547        pub span: Span,
548        pub def_descr: &'static str,
549    }
550
551    #[derive(Diagnostic)]
552    #[diag("call to {$callee} may panic: {$reason}", code = E0999)]
553    pub(super) struct PanicError {
554        #[primary_span]
555        pub(super) span: Span,
556        pub(super) callee: String,
557        pub(super) reason: String,
558    }
559}