use std::{
error::Error as StdError,
future::Future,
process::{Command, Stdio},
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
time::{Duration, Instant},
};
use crate::proto::api::ProverServiceClient;
use async_trait::async_trait;
use proto::api::ReadyRequest;
use reqwest::{Request, Response};
use serde::{Deserialize, Serialize};
use sp1_core_machine::{io::SP1Stdin, reduce::SP1ReduceProof, utils::SP1CoreProverError};
use sp1_prover::{
InnerSC, OuterSC, SP1CoreProof, SP1ProvingKey, SP1RecursionProverError, SP1VerifyingKey,
};
use tokio::task::block_in_place;
use twirp::{
async_trait,
reqwest::{self},
url::Url,
Client, ClientError, Middleware, Next,
};
#[rustfmt::skip]
pub mod proto {
pub mod api;
}
pub struct SP1CudaProver {
client: Client,
managed_container: Option<CudaProverContainer>,
}
pub struct CudaProverContainer {
name: String,
cleaned_up: Arc<AtomicBool>,
}
#[derive(Serialize, Deserialize)]
pub struct SetupRequestPayload {
pub elf: Vec<u8>,
}
#[derive(Serialize, Deserialize)]
pub struct SetupResponsePayload {
pub pk: SP1ProvingKey,
pub vk: SP1VerifyingKey,
}
#[derive(Serialize, Deserialize)]
pub struct ProveCoreRequestPayload {
pub stdin: SP1Stdin,
}
#[derive(Serialize, Deserialize)]
pub struct CompressRequestPayload {
pub vk: SP1VerifyingKey,
pub proof: SP1CoreProof,
pub deferred_proofs: Vec<SP1ReduceProof<InnerSC>>,
}
#[derive(Serialize, Deserialize)]
pub struct ShrinkRequestPayload {
pub reduced_proof: SP1ReduceProof<InnerSC>,
}
#[derive(Serialize, Deserialize)]
pub struct WrapRequestPayload {
pub reduced_proof: SP1ReduceProof<InnerSC>,
}
impl SP1CudaProver {
pub fn new(moongate_endpoint: Option<String>) -> Result<Self, Box<dyn StdError>> {
let reqwest_middlewares = vec![Box::new(LoggingMiddleware) as Box<dyn Middleware>];
let prover = match moongate_endpoint {
Some(moongate_endpoint) => {
let client = Client::new(
Url::parse(&moongate_endpoint).expect("failed to parse url"),
reqwest::Client::new(),
reqwest_middlewares,
)
.expect("failed to create client");
SP1CudaProver { client, managed_container: None }
}
None => Self::start_moongate_server(reqwest_middlewares)?,
};
let timeout = Duration::from_secs(300);
let start_time = Instant::now();
block_on(async {
tracing::info!("waiting for proving server to be ready");
loop {
if start_time.elapsed() > timeout {
return Err("Timeout: proving server did not become ready within 60 seconds. Please check your Docker container and network settings.".to_string());
}
let request = ReadyRequest {};
match prover.client.ready(request).await {
Ok(response) if response.ready => {
tracing::info!("proving server is ready");
break;
}
Ok(_) => {
tracing::info!("proving server is not ready, retrying...");
}
Err(e) => {
tracing::warn!("Error checking server readiness: {}", e);
}
}
tokio::time::sleep(Duration::from_secs(2)).await;
}
Ok(())
})?;
Ok(prover)
}
fn check_docker_availability() -> Result<bool, Box<dyn std::error::Error>> {
match Command::new("docker").arg("version").output() {
Ok(output) => Ok(output.status.success()),
Err(_) => Ok(false),
}
}
fn start_moongate_server(
reqwest_middlewares: Vec<Box<dyn Middleware>>,
) -> Result<SP1CudaProver, Box<dyn StdError>> {
let container_name = "sp1-gpu";
let image_name = std::env::var("SP1_GPU_IMAGE")
.unwrap_or_else(|_| "public.ecr.aws/succinct-labs/moongate:v4.1.0".to_string());
let cleaned_up = Arc::new(AtomicBool::new(false));
let cleanup_name = container_name;
let cleanup_flag = cleaned_up.clone();
if !Self::check_docker_availability()? {
return Err("Docker is not available or you don't have the necessary permissions. Please ensure Docker is installed and you are part of the docker group.".into());
}
if let Err(e) = Command::new("docker").args(["pull", &image_name]).output() {
return Err(format!("Failed to pull Docker image: {}. Please check your internet connection and Docker permissions.", e).into());
}
let rust_log_level = std::env::var("RUST_LOG").unwrap_or_else(|_| "none".to_string());
Command::new("docker")
.args([
"run",
"-e",
&format!("RUST_LOG={}", rust_log_level),
"-p",
"3000:3000",
"--rm",
"--gpus",
"all",
"--name",
container_name,
&image_name,
])
.stdout(Stdio::inherit())
.stderr(Stdio::inherit())
.spawn()
.map_err(|e| format!("Failed to start Docker container: {}. Please check your Docker installation and permissions.", e))?;
ctrlc::set_handler(move || {
tracing::debug!("received Ctrl+C, cleaning up...");
if !cleanup_flag.load(Ordering::SeqCst) {
cleanup_container(cleanup_name);
cleanup_flag.store(true, Ordering::SeqCst);
}
std::process::exit(0);
})
.unwrap();
std::thread::sleep(Duration::from_secs(2));
let client = Client::new(
Url::parse("http://localhost:3000/twirp/").expect("failed to parse url"),
reqwest::Client::new(),
reqwest_middlewares,
)
.expect("failed to create client");
Ok(SP1CudaProver {
client,
managed_container: Some(CudaProverContainer {
name: container_name.to_string(),
cleaned_up: cleaned_up.clone(),
}),
})
}
pub fn setup(&self, elf: &[u8]) -> Result<(SP1ProvingKey, SP1VerifyingKey), Box<dyn StdError>> {
let payload = SetupRequestPayload { elf: elf.to_vec() };
let request =
crate::proto::api::SetupRequest { data: bincode::serialize(&payload).unwrap() };
let response = block_on(async { self.client.setup(request).await }).unwrap();
let payload: SetupResponsePayload = bincode::deserialize(&response.result).unwrap();
Ok((payload.pk, payload.vk))
}
pub fn prove_core(&self, stdin: &SP1Stdin) -> Result<SP1CoreProof, SP1CoreProverError> {
let payload = ProveCoreRequestPayload { stdin: stdin.clone() };
let request =
crate::proto::api::ProveCoreRequest { data: bincode::serialize(&payload).unwrap() };
let response = block_on(async { self.client.prove_core(request).await }).unwrap();
let proof: SP1CoreProof = bincode::deserialize(&response.result).unwrap();
Ok(proof)
}
pub fn compress(
&self,
vk: &SP1VerifyingKey,
proof: SP1CoreProof,
deferred_proofs: Vec<SP1ReduceProof<InnerSC>>,
) -> Result<SP1ReduceProof<InnerSC>, SP1RecursionProverError> {
let payload = CompressRequestPayload { vk: vk.clone(), proof, deferred_proofs };
let request =
crate::proto::api::CompressRequest { data: bincode::serialize(&payload).unwrap() };
let response = block_on(async { self.client.compress(request).await }).unwrap();
let proof: SP1ReduceProof<InnerSC> = bincode::deserialize(&response.result).unwrap();
Ok(proof)
}
pub fn shrink(
&self,
reduced_proof: SP1ReduceProof<InnerSC>,
) -> Result<SP1ReduceProof<InnerSC>, SP1RecursionProverError> {
let payload = ShrinkRequestPayload { reduced_proof: reduced_proof.clone() };
let request =
crate::proto::api::ShrinkRequest { data: bincode::serialize(&payload).unwrap() };
let response = block_on(async { self.client.shrink(request).await }).unwrap();
let proof: SP1ReduceProof<InnerSC> = bincode::deserialize(&response.result).unwrap();
Ok(proof)
}
pub fn wrap_bn254(
&self,
reduced_proof: SP1ReduceProof<InnerSC>,
) -> Result<SP1ReduceProof<OuterSC>, SP1RecursionProverError> {
let payload = WrapRequestPayload { reduced_proof: reduced_proof.clone() };
let request =
crate::proto::api::WrapRequest { data: bincode::serialize(&payload).unwrap() };
let response = block_on(async { self.client.wrap(request).await }).unwrap();
let proof: SP1ReduceProof<OuterSC> = bincode::deserialize(&response.result).unwrap();
Ok(proof)
}
}
impl Default for SP1CudaProver {
fn default() -> Self {
Self::new(None).expect("Failed to create SP1CudaProver")
}
}
impl Drop for SP1CudaProver {
fn drop(&mut self) {
if let Some(container) = &self.managed_container {
if !container.cleaned_up.load(Ordering::SeqCst) {
tracing::debug!("dropping SP1ProverClient, cleaning up...");
cleanup_container(&container.name);
container.cleaned_up.store(true, Ordering::SeqCst);
}
}
}
}
fn cleanup_container(container_name: &str) {
if let Err(e) = Command::new("docker").args(["rm", "-f", container_name]).output() {
eprintln!(
"Failed to remove container: {}. You may need to manually remove it using 'docker rm -f {}'",
e, container_name
);
}
}
pub fn block_on<T>(fut: impl Future<Output = T>) -> T {
if let Ok(handle) = tokio::runtime::Handle::try_current() {
block_in_place(|| handle.block_on(fut))
} else {
let rt = tokio::runtime::Runtime::new().expect("Failed to create a new runtime");
rt.block_on(fut)
}
}
struct LoggingMiddleware;
pub type Result<T, E = ClientError> = std::result::Result<T, E>;
#[async_trait]
impl Middleware for LoggingMiddleware {
async fn handle(&self, req: Request, next: Next<'_>) -> Result<Response> {
let response = next.run(req).await;
match response {
Ok(response) => {
tracing::info!("{:?}", response);
Ok(response)
}
Err(e) => Err(e),
}
}
}