use super::smtlib_script::SmtLibScript;
use super::solver::{Decision, DecisionWithModel, LocalSolver, Solver, SolverError};
use miette::Diagnostic;
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use thiserror::Error;
use tokio::io::AsyncWrite;
use tokio::sync::{Mutex, OwnedSemaphorePermit, Semaphore};
#[derive(Debug, Diagnostic, Error)]
pub enum SolverPoolError {
#[error("solver pool semaphore closed unexpectedly")]
SemaphoreClosed,
#[error("timeout waiting for solver from pool")]
AcquireTimeout,
#[error(transparent)]
#[diagnostic(transparent)]
Solver(#[from] SolverError),
}
#[derive(Debug, Default)]
struct FailedWriter;
impl FailedWriter {
fn error() -> io::Error {
io::Error::other(SolverError::SolverMarkedFailed)
}
}
impl AsyncWrite for FailedWriter {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &[u8],
) -> Poll<io::Result<usize>> {
Poll::Ready(Err(Self::error()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Err(Self::error()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Err(Self::error()))
}
}
#[derive(Clone, Debug)]
pub struct SolverPoolConfig {
pub min_solvers: usize,
pub max_solvers: usize,
pub acquire_timeout: Option<Duration>,
}
impl Default for SolverPoolConfig {
fn default() -> Self {
Self {
min_solvers: 1,
max_solvers: 4,
acquire_timeout: None,
}
}
}
#[derive(Debug)]
pub struct SolverPool {
available: Arc<Mutex<Vec<LocalSolver>>>,
semaphore: Arc<Semaphore>,
config: SolverPoolConfig,
}
impl SolverPool {
pub async fn new(config: SolverPoolConfig) -> Result<Self, SolverError> {
let semaphore = Arc::new(Semaphore::new(config.max_solvers));
let mut solvers = Vec::with_capacity(config.max_solvers);
for _ in 0..config.min_solvers {
solvers.push(LocalSolver::cvc5()?);
}
Ok(Self {
available: Arc::new(Mutex::new(solvers)),
semaphore,
config,
})
}
pub async fn acquire(&self) -> Result<PooledSolver, SolverPoolError> {
let permit = match self.config.acquire_timeout {
Some(timeout) => {
match tokio::time::timeout(timeout, Arc::clone(&self.semaphore).acquire_owned())
.await
{
Ok(Ok(permit)) => permit,
Ok(Err(_)) => return Err(SolverPoolError::SemaphoreClosed),
Err(_) => return Err(SolverPoolError::AcquireTimeout),
}
}
None => Arc::clone(&self.semaphore)
.acquire_owned()
.await
.map_err(|_| SolverPoolError::SemaphoreClosed)?,
};
let solver = {
let mut available = self.available.lock().await;
available.pop()
};
let solver = match solver {
Some(s) => s,
None => {
LocalSolver::cvc5()?
}
};
Ok(PooledSolver {
solver: Some(solver),
pool: Arc::clone(&self.available),
permit: Some(permit),
failed_writer: FailedWriter,
})
}
pub async fn available_count(&self) -> usize {
self.available.lock().await.len()
}
pub fn permits_available(&self) -> usize {
self.semaphore.available_permits()
}
}
#[derive(Debug)]
pub struct PooledSolver {
solver: Option<LocalSolver>,
pool: Arc<Mutex<Vec<LocalSolver>>>,
permit: Option<OwnedSemaphorePermit>,
failed_writer: FailedWriter,
}
impl PooledSolver {
pub fn mark_failed(&mut self) {
self.solver = None;
}
}
impl Solver for PooledSolver {
fn smtlib_input(&mut self) -> &mut (dyn tokio::io::AsyncWrite + Unpin + Send) {
match &mut self.solver {
Some(s) => s.smtlib_input(),
None => &mut self.failed_writer,
}
}
async fn enable_models(&mut self) -> Result<(), SolverError> {
let solver = self
.solver
.as_mut()
.ok_or(SolverError::SolverMarkedFailed)?;
solver.enable_models().await
}
async fn check_sat(&mut self) -> Result<Decision, SolverError> {
let solver = self
.solver
.as_mut()
.ok_or(SolverError::SolverMarkedFailed)?;
solver.check_sat().await
}
async fn check_sat_with_model(&mut self) -> Result<DecisionWithModel, SolverError> {
let solver = self
.solver
.as_mut()
.ok_or(SolverError::SolverMarkedFailed)?;
solver.check_sat_with_model().await
}
}
impl Drop for PooledSolver {
fn drop(&mut self) {
let pool = Arc::clone(&self.pool);
let solver = self.solver.take();
let permit = self.permit.take();
tokio::spawn(async move {
if let Some(mut solver) = solver {
if solver.smtlib_input().reset().await.is_ok() {
pool.lock().await.push(solver);
}
}
drop(permit);
});
}
}
#[cfg(test)]
mod test {
use super::*;
use cool_asserts::assert_matches;
async fn poll_until<F, Fut>(mut condition: F)
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = bool>,
{
for _ in 0..10 {
if condition().await {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
panic!("condition not met after 100ms of polling");
}
#[tokio::test]
async fn test_pool_creation() {
let pool = SolverPool::new(SolverPoolConfig {
min_solvers: 2,
max_solvers: 4,
acquire_timeout: None,
})
.await
.unwrap();
assert_eq!(pool.available_count().await, 2);
assert_eq!(pool.permits_available(), 4);
}
#[tokio::test]
async fn test_acquire_release() {
let pool = SolverPool::new(SolverPoolConfig {
min_solvers: 1,
max_solvers: 4,
acquire_timeout: None,
})
.await
.unwrap();
assert_eq!(pool.available_count().await, 1);
{
let mut solver = pool.acquire().await.unwrap();
assert_eq!(pool.available_count().await, 0);
let decision = solver.check_sat().await.unwrap();
assert_eq!(decision, Decision::Sat);
}
poll_until(|| async { pool.available_count().await == 1 }).await;
}
#[tokio::test]
async fn test_solver_reuse() {
let pool = SolverPool::new(SolverPoolConfig {
min_solvers: 1,
max_solvers: 1,
acquire_timeout: None,
})
.await
.unwrap();
{
let mut solver = pool.acquire().await.unwrap();
solver.smtlib_input().assert("false").await.unwrap();
let decision = solver.check_sat().await.unwrap();
assert_eq!(decision, Decision::Unsat);
}
poll_until(|| async { pool.available_count().await == 1 }).await;
{
let mut solver = pool.acquire().await.unwrap();
solver.enable_models().await.unwrap();
solver.smtlib_input().assert("true").await.unwrap();
let decision = solver.check_sat_with_model().await.unwrap();
assert_matches!(decision, DecisionWithModel::Sat { .. });
}
}
#[tokio::test]
async fn test_concurrent_acquire() {
let pool = Arc::new(
SolverPool::new(SolverPoolConfig {
min_solvers: 2,
max_solvers: 4,
acquire_timeout: None,
})
.await
.unwrap(),
);
let handles: Vec<_> = (0..4)
.map(|_| {
let pool = Arc::clone(&pool);
tokio::spawn(async move {
let mut solver = pool.acquire().await.unwrap();
let decision = solver.check_sat().await.unwrap();
assert_eq!(decision, Decision::Sat);
})
})
.collect();
for handle in handles {
handle.await.unwrap();
}
}
#[tokio::test]
async fn test_pool_exhaustion_blocks() {
let pool = Arc::new(
SolverPool::new(SolverPoolConfig {
min_solvers: 1,
max_solvers: 1,
acquire_timeout: Some(Duration::from_millis(100)),
})
.await
.unwrap(),
);
let _solver = pool.acquire().await.unwrap();
let result = pool.acquire().await;
assert_matches!(result, Err(SolverPoolError::AcquireTimeout));
}
#[tokio::test]
async fn test_failed_solver_discarded() {
let pool = SolverPool::new(SolverPoolConfig {
min_solvers: 1,
max_solvers: 1,
acquire_timeout: None,
})
.await
.unwrap();
{
let mut solver = pool.acquire().await.unwrap();
solver.smtlib_input().assert("tomato").await.unwrap();
let result = solver.check_sat().await;
assert_matches!(result, Err(SolverError::Solver(_)));
solver.mark_failed();
}
poll_until(|| async { pool.permits_available() == 1 }).await;
assert_eq!(pool.available_count().await, 0);
let mut solver = pool.acquire().await.unwrap();
let decision = solver.check_sat().await.unwrap();
assert_eq!(decision, Decision::Sat);
}
#[tokio::test]
async fn test_methods_after_mark_failed() {
use tokio::io::AsyncWriteExt;
let pool = SolverPool::new(SolverPoolConfig::default()).await.unwrap();
let mut solver = pool.acquire().await.unwrap();
solver.mark_failed();
let result = solver.check_sat().await;
assert_matches!(result, Err(SolverError::SolverMarkedFailed));
let result = solver.check_sat_with_model().await;
assert_matches!(result, Err(SolverError::SolverMarkedFailed));
let result = solver.smtlib_input().write_all(b"test").await;
assert!(result.is_err());
let result = solver.smtlib_input().flush().await;
assert!(result.is_err());
let result = solver.smtlib_input().shutdown().await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_default_config() {
let config = SolverPoolConfig::default();
assert_eq!(config.min_solvers, 1);
assert_eq!(config.max_solvers, 4);
assert!(config.acquire_timeout.is_none());
}
#[tokio::test]
async fn test_pooled_solver_with_symcompiler() {
use crate::symcc::SymCompiler;
let pool = SolverPool::new(SolverPoolConfig::default()).await.unwrap();
let solver = pool.acquire().await.unwrap();
let mut compiler = SymCompiler::new(solver);
compiler
.solver_mut()
.smtlib_input()
.set_logic("ALL")
.await
.unwrap();
compiler
.solver_mut()
.smtlib_input()
.assert("(= 1 2)")
.await
.unwrap();
let decision = compiler.solver_mut().check_sat().await.unwrap();
assert_eq!(decision, Decision::Unsat);
}
}