use std::{
fmt::Debug,
hash::Hash,
sync::{atomic::AtomicI32, Arc, Mutex},
};
use teloxide::{
dispatching::UpdateHandler,
error_handlers::ErrorHandler,
prelude::*,
stop::mk_stop_token,
types::{MaybeInaccessibleMessage, Me, UpdateKind},
};
pub use crate::utils::DistributionKey;
use crate::{
dataset::{IntoUpdate, MockMe},
server,
server::ServerManager,
state::State,
utils::default_distribution_function,
};
pub struct MockBot<Err, Key> {
pub bot: Bot,
handler_tree: Arc<UpdateHandler<Err>>,
pub updates: Vec<Update>,
pub me: Me,
pub dependencies: DependencyMap,
distribution_f: fn(&Update) -> Option<Key>,
error_handler: Arc<dyn ErrorHandler<Err> + Send + Sync>,
current_update_id: AtomicI32,
state: Arc<Mutex<State>>,
#[allow(dead_code)]
server: ServerManager,
api_url: url::Url,
}
impl<Err> MockBot<Err, DistributionKey>
where
Err: Debug + Send + Sync + 'static,
{
pub async fn new<T>(update: T, handler_tree: UpdateHandler<Err>) -> Self
where
T: IntoUpdate,
Err: Debug,
{
let _ = pretty_env_logger::try_init();
let token = "1234567890:QWERTYUIOPASDFGHJKLZXCVBNMQWERTYUIO";
let bot = Bot::new(token);
let current_update_id = AtomicI32::new(42);
let state = Arc::new(Mutex::new(State::default()));
let me = MockMe::new().build();
let server = ServerManager::start(me.clone(), state.clone())
.await
.expect("Failed to start mock server");
let api_url = url::Url::parse(&format!("http://127.0.0.1:{}", server.port))
.expect("Failed to parse API URL");
Self {
bot,
me,
updates: update.into_update(¤t_update_id),
handler_tree: Arc::new(handler_tree), dependencies: DependencyMap::new(),
error_handler: LoggingErrorHandler::new(),
distribution_f: default_distribution_function,
current_update_id,
state,
server,
api_url,
}
}
}
impl<Err, Key> MockBot<Err, Key>
where
Err: Debug + Send + Sync + 'static,
Key: Hash + Eq + Clone + Send + 'static,
{
pub async fn new_with_distribution_function<T>(
update: T,
handler_tree: UpdateHandler<Err>,
f: fn(&Update) -> Option<Key>,
) -> Self
where
T: IntoUpdate,
Err: Debug,
{
let _ = pretty_env_logger::try_init();
let token = "1234567890:QWERTYUIOPASDFGHJKLZXCVBNMQWERTYUIO";
let bot = Bot::new(token);
let current_update_id = AtomicI32::new(42);
let state = Arc::new(Mutex::new(State::default()));
let me = MockMe::new().build();
let server = ServerManager::start(me.clone(), state.clone())
.await
.expect("Failed to start mock server");
let api_url = url::Url::parse(&format!("http://127.0.0.1:{}", server.port))
.expect("Failed to parse API URL");
Self {
bot,
me,
updates: update.into_update(¤t_update_id),
handler_tree: Arc::new(handler_tree),
dependencies: DependencyMap::new(),
error_handler: LoggingErrorHandler::new(),
distribution_f: f,
current_update_id,
state,
server,
api_url,
}
}
pub fn dependencies(&mut self, deps: DependencyMap) {
self.dependencies = deps;
}
pub fn me(&mut self, me: MockMe) {
self.me = me.build();
}
pub fn update<T: IntoUpdate>(&mut self, update: T) {
self.updates = update.into_update(&self.current_update_id);
}
pub fn error_handler(&mut self, handler: Arc<dyn ErrorHandler<Err> + Send + Sync>) {
self.error_handler = handler;
}
pub fn api_url(&self) -> &url::Url {
&self.api_url
}
fn insert_updates(&self, updates: &mut [Update]) {
let mut state = self.state.lock().unwrap();
for update in updates.iter_mut() {
match &mut update.kind {
UpdateKind::Message(ref mut message) => {
state.add_message(message);
}
UpdateKind::EditedMessage(ref mut message) => {
state.edit_message(message);
}
UpdateKind::CallbackQuery(ref mut callback) => {
if let Some(MaybeInaccessibleMessage::Regular(ref mut message)) =
callback.message
{
state.add_message(message);
}
}
_ => {}
}
}
}
pub async fn dispatch(&mut self) {
self.state.lock().unwrap().reset();
let mut updates = self.updates.clone();
self.insert_updates(&mut updates);
let bot = self.bot.clone().set_api_url(self.api_url.clone());
let handler_tree = Arc::clone(&self.handler_tree);
let deps = self.dependencies.clone();
let distribution_f = self.distribution_f;
let error_handler = self.error_handler.clone();
let handle = tokio::task::spawn(async move {
Dispatcher::builder(bot, (*handler_tree).clone())
.dependencies(deps)
.distribution_function(distribution_f)
.error_handler(error_handler)
.build()
.dispatch_with_listener(
SingleUpdateListener::new(updates),
LoggingErrorHandler::new(),
)
.await;
});
handle.await.expect("Dispatch task panicked!");
}
pub fn get_responses(&self) -> server::Responses {
self.state.lock().unwrap().responses.clone()
}
}
struct SingleUpdateListener {
updates: Vec<Update>,
}
impl SingleUpdateListener {
fn new(updates: Vec<Update>) -> Self {
Self { updates }
}
}
impl teloxide::update_listeners::UpdateListener for SingleUpdateListener {
type Err = std::convert::Infallible;
fn stop_token(&mut self) -> teloxide::stop::StopToken {
let (token, _flag) = mk_stop_token();
token
}
}
impl<'a> teloxide::update_listeners::AsUpdateStream<'a> for SingleUpdateListener {
type StreamErr = std::convert::Infallible;
type Stream = SingleUpdateStream;
fn as_stream(&'a mut self) -> Self::Stream {
SingleUpdateStream {
updates: std::mem::take(&mut self.updates).into(),
}
}
}
struct SingleUpdateStream {
updates: std::collections::VecDeque<Update>,
}
impl futures_util::Stream for SingleUpdateStream {
type Item = Result<Update, std::convert::Infallible>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
match self.updates.pop_front() {
Some(update) => std::task::Poll::Ready(Some(Ok(update))),
None => std::task::Poll::Ready(None),
}
}
}