use opcda_bridge_proto::bridge::bridge_server::{Bridge, BridgeServer};
use opcda_bridge_proto::bridge::{
BrowseRequest, BrowseResponse, ListServersRequest, ListServersResponse, ReadRequest,
ReadResponse, WriteRequest, WriteResponse,
};
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::sync::{Notify, mpsc};
use tokio_stream::wrappers::ReceiverStream;
use tonic::transport::Server;
use tonic::{Request, Response, Status};
#[derive(Default)]
pub(crate) struct MockBridgeService {
pub(crate) list_servers_response: ListServersResponse,
pub(crate) list_servers_error: Option<Status>,
pub(crate) browse_responses: Vec<BrowseResponse>,
pub(crate) browse_initial_error: Option<Status>,
pub(crate) browse_stream_error: Option<Status>,
pub(crate) browse_send_failure: Arc<Notify>,
pub(crate) server_shutdown: Arc<Notify>,
pub(crate) server_stopped: Arc<Notify>,
pub(crate) read_response: ReadResponse,
pub(crate) read_error: Option<Status>,
pub(crate) write_response: WriteResponse,
pub(crate) write_error: Option<Status>,
}
#[tonic::async_trait]
impl Bridge for MockBridgeService {
async fn list_servers(
&self,
_request: Request<ListServersRequest>,
) -> Result<Response<ListServersResponse>, Status> {
if let Some(status) = self.list_servers_error.clone() {
return Err(status);
}
Ok(Response::new(self.list_servers_response.clone()))
}
type BrowseStream = ReceiverStream<Result<BrowseResponse, Status>>;
async fn browse(
&self,
_request: Request<BrowseRequest>,
) -> Result<Response<Self::BrowseStream>, Status> {
if let Some(status) = self.browse_initial_error.clone() {
return Err(status);
}
let (tx, rx) = mpsc::channel(4);
let items = self.browse_responses.clone();
let stream_error = self.browse_stream_error.clone();
let browse_send_failure = Arc::clone(&self.browse_send_failure);
tokio::spawn(async move {
for item in items {
if tx.send(Ok(item)).await.is_err() {
browse_send_failure.notify_one();
break;
}
}
if let Some(status) = stream_error {
let _ = tx.send(Err(status)).await;
}
});
Ok(Response::new(ReceiverStream::new(rx)))
}
async fn read(&self, _request: Request<ReadRequest>) -> Result<Response<ReadResponse>, Status> {
if let Some(status) = self.read_error.clone() {
return Err(status);
}
Ok(Response::new(self.read_response.clone()))
}
async fn write(
&self,
_request: Request<WriteRequest>,
) -> Result<Response<WriteResponse>, Status> {
if let Some(status) = self.write_error.clone() {
return Err(status);
}
Ok(Response::new(self.write_response.clone()))
}
}
pub(crate) async fn start_mock_server(service: MockBridgeService) -> String {
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let listener = tokio::net::TcpListener::bind(addr).await.unwrap();
let port = listener.local_addr().unwrap().port();
let server_shutdown = Arc::clone(&service.server_shutdown);
let server_stopped = Arc::clone(&service.server_stopped);
tokio::spawn(async move {
Server::builder()
.add_service(BridgeServer::new(service))
.serve_with_incoming_shutdown(
tokio_stream::wrappers::TcpListenerStream::new(listener),
server_shutdown.notified(),
)
.await
.unwrap();
server_stopped.notify_one();
});
format!("127.0.0.1:{port}")
}