Skip to main content

cargo_airbender/commands/
prove.rs

1use crate::cli::{ProveArgs, ProverBackendArg, ProverLevelArg};
2use crate::error::{CliError, Result};
3use crate::input;
4use crate::ui;
5use airbender_host::Prover;
6
7pub fn run(args: ProveArgs) -> Result<()> {
8    let input_words = input::parse_input_words(&args.input)?;
9    let security = args.security;
10
11    let prove_result = match args.backend {
12        ProverBackendArg::Dev => {
13            if args.threads.is_some() {
14                tracing::warn!("ignoring `--threads` for dev backend");
15            }
16            if args.ram_bound.is_some() {
17                tracing::warn!("ignoring `--ram-bound` for dev backend");
18            }
19            if args.level != ProverLevelArg::RecursionUnified {
20                tracing::warn!("ignoring `--level` for dev backend");
21            }
22
23            let security = security.into();
24            let prover = airbender_host::DevProverBuilder::new(&args.app_bin)
25                .with_security(security)
26                .maybe_cycles(args.cycles)
27                .build()
28                .map_err(|err| {
29                CliError::with_source(
30                    format!(
31                        "failed to initialize dev prover for `{}`",
32                        args.app_bin.display()
33                    ),
34                    err,
35                )
36            })?;
37
38            prover.prove(&input_words)
39        }
40        ProverBackendArg::Gpu => {
41            if args.cycles.is_some() {
42                tracing::warn!("ignoring `--cycles` for gpu backend");
43            }
44            if args.ram_bound.is_some() {
45                tracing::warn!("ignoring `--ram-bound` for gpu backend");
46            }
47
48            #[cfg(feature = "gpu-prover")]
49            {
50                let level = as_host_level(args.level);
51                let security = security.into();
52                let prover = airbender_host::GpuProverBuilder::new(&args.app_bin)
53                    .with_level(level)
54                    .with_security(security)
55                    .with_config(
56                        airbender_host::GpuProverConfig::default()
57                            .maybe_worker_threads(args.threads),
58                    )
59                    .build()
60                    .map_err(|err| {
61                    CliError::with_source(
62                        format!(
63                            "failed to initialize GPU prover for `{}`",
64                            args.app_bin.display()
65                        ),
66                        err,
67                    )
68                })?;
69
70                prover.prove(&input_words)
71            }
72
73            #[cfg(not(feature = "gpu-prover"))]
74            {
75                return Err(CliError::new(
76                    "GPU backend requires GPU support in `cargo-airbender`",
77                )
78                .with_hint(
79                    "rebuild `cargo-airbender` with default features or pass `--features gpu-prover` to use `--backend gpu`",
80                ));
81            }
82        }
83        ProverBackendArg::Cpu => {
84            let level = as_host_level(args.level);
85            let security = security.into();
86            if level != airbender_host::ProverLevel::Base {
87                return Err(
88                    CliError::new("CPU backend currently supports only `--level base`")
89                        .with_hint("use `--backend gpu` for recursion levels"),
90                );
91            }
92
93            let prover = airbender_host::CpuProverBuilder::new(&args.app_bin)
94                .with_security(security)
95                .maybe_worker_threads(args.threads)
96                .maybe_cycles(args.cycles)
97                .maybe_ram_bound(args.ram_bound)
98                .build()
99                .map_err(|err| {
100                CliError::with_source(
101                    format!(
102                        "failed to initialize CPU prover for `{}`",
103                        args.app_bin.display()
104                    ),
105                    err,
106                )
107            })?;
108
109            prover.prove(&input_words)
110        }
111    }
112    .map_err(|err| {
113        CliError::with_source(
114            format!("failed to generate proof for `{}`", args.app_bin.display()),
115            err,
116        )
117        .with_hint("set `RUST_LOG=info` to inspect prover backend logs")
118    })?;
119
120    tracing::info!("{}", prove_result.proof.debug_info());
121
122    let encoded = bincode::serde::encode_to_vec(&prove_result.proof, bincode::config::standard())
123        .map_err(|err| CliError::with_source("failed to encode proof", err))?;
124    std::fs::write(&args.output, encoded).map_err(|err| {
125        CliError::with_source(
126            format!("failed to write proof to `{}`", args.output.display()),
127            err,
128        )
129    })?;
130
131    ui::success("proof generated");
132    ui::field("backend", backend_name(args.backend));
133    ui::field("level", proof_level(args.backend, args.level));
134    ui::field("security", security);
135    ui::field("cycles", prove_result.cycles);
136    ui::field("output", args.output.display());
137
138    Ok(())
139}
140
141fn backend_name(backend: ProverBackendArg) -> &'static str {
142    match backend {
143        ProverBackendArg::Dev => "dev",
144        ProverBackendArg::Cpu => "cpu",
145        ProverBackendArg::Gpu => "gpu",
146    }
147}
148
149fn proof_level(backend: ProverBackendArg, level: ProverLevelArg) -> &'static str {
150    match backend {
151        ProverBackendArg::Dev => "dev",
152        ProverBackendArg::Cpu | ProverBackendArg::Gpu => level_name(level),
153    }
154}
155
156fn level_name(level: ProverLevelArg) -> &'static str {
157    match level {
158        ProverLevelArg::Base => "base",
159        ProverLevelArg::RecursionUnrolled => "recursion-unrolled",
160        ProverLevelArg::RecursionUnified => "recursion-unified",
161    }
162}
163
164fn as_host_level(level: ProverLevelArg) -> airbender_host::ProverLevel {
165    match level {
166        ProverLevelArg::Base => airbender_host::ProverLevel::Base,
167        ProverLevelArg::RecursionUnrolled => airbender_host::ProverLevel::RecursionUnrolled,
168        ProverLevelArg::RecursionUnified => airbender_host::ProverLevel::RecursionUnified,
169    }
170}