use crate::mock_server::hyper::run_server;
use crate::mock_set::ActiveMockSet;
use crate::{mock::Mock, verification::VerificationOutcome, Request};
use std::net::{SocketAddr, TcpListener, TcpStream};
use std::sync::Arc;
use tokio::sync::{Mutex, RwLock};
use tokio::task::LocalSet;
pub(crate) struct BareMockServer {
mock_set: Arc<RwLock<ActiveMockSet>>,
received_requests: Option<Arc<Mutex<Vec<Request>>>>,
server_address: SocketAddr,
_shutdown_trigger: tokio::sync::oneshot::Sender<()>,
}
impl BareMockServer {
pub(super) async fn start(listener: TcpListener, request_recording: RequestRecording) -> Self {
let (shutdown_trigger, shutdown_receiver) = tokio::sync::oneshot::channel();
let mock_set = Arc::new(RwLock::new(ActiveMockSet::new()));
let received_requests = match request_recording {
RequestRecording::Enabled => Some(Arc::new(Mutex::new(Vec::new()))),
RequestRecording::Disabled => None,
};
let server_address = listener
.local_addr()
.expect("Failed to get server address.");
let server_mock_set = mock_set.clone();
let server_received_requests = received_requests.clone();
std::thread::spawn(move || {
let server_future = run_server(
listener,
server_mock_set,
server_received_requests,
shutdown_receiver,
);
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("Cannot build local tokio runtime");
LocalSet::new().block_on(&runtime, server_future)
});
for _ in 0..40 {
if TcpStream::connect_timeout(&server_address, std::time::Duration::from_millis(25))
.is_ok()
{
break;
}
futures_timer::Delay::new(std::time::Duration::from_millis(25)).await;
}
Self {
mock_set,
received_requests,
server_address,
_shutdown_trigger: shutdown_trigger,
}
}
pub(crate) async fn register(&self, mock: Mock) {
self.mock_set.write().await.register(mock);
}
pub(crate) async fn reset(&self) {
self.mock_set.write().await.reset();
if let Some(received_requests) = &self.received_requests {
received_requests.lock().await.clear();
}
}
pub(crate) async fn verify(&self) -> VerificationOutcome {
let mock_set = self.mock_set.read().await;
mock_set.verify()
}
pub(crate) fn uri(&self) -> String {
format!("http://{}", self.server_address)
}
pub(crate) fn address(&self) -> &SocketAddr {
&self.server_address
}
pub(crate) async fn received_requests(&self) -> Option<Vec<Request>> {
if let Some(received_requests) = &self.received_requests {
Some(received_requests.lock().await.clone())
} else {
None
}
}
}
pub(super) enum RequestRecording {
Enabled,
Disabled,
}