Skip to content
Draft
Next Next commit
Initial FSM PoC
  • Loading branch information
Bibo-Joshi committed Feb 4, 2025
commit 0c06ba0a9053d5e85a7df6aa71ef464b78326c55
195 changes: 195 additions & 0 deletions examples/fsmbot.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,195 @@
#!/usr/bin/env python
# pylint: disable=unused-argument
# This program is dedicated to the public domain under the CC0 license.
"""Simple state machine to handle user support.
One admin is supported. The admin can have one active conversation at a time. Other users
are put on hold until the admin finishes the current conversation.
In each conversation, the admin and the user take turns to send messages.
"""
import logging
from typing import Optional

from telegram import Update
from telegram.ext import (
Application,
CommandHandler,
ContextTypes,
FiniteStateMachine,
MessageHandler,
State,
filters,
)

# Enable logging
logging.basicConfig(
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", level=logging.DEBUG
)
logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("httpcore").setLevel(logging.WARNING)
logging.getLogger("telegram").setLevel(logging.WARNING)
logging.getLogger("telegram.ext.Application").setLevel(logging.DEBUG)

logger = logging.getLogger(__name__)


class UserSupportMachine(FiniteStateMachine[Optional[int]]):

HOLD = State("HOLD")
WELCOMING = State("WELCOMING")
WAITING_FOR_REPLY = State("WAITING_FOR_REPLY")
WRITING = State("WRITING")

def __init__(self, admin_id: int):
self.admin_id = admin_id
self._states: dict[int, State] = {}
super().__init__()

def _get_admin_state(self) -> State:
return self._states.get(self.admin_id, State.IDLE)

def get_active_key_state(self, update: object) -> tuple[Optional[int], State]:
if not isinstance(update, Update):
return None, State.IDLE
if not (user := update.effective_user):
return None, State.IDLE

# Admin is easy - just return the state
admin_state = self._get_admin_state()
if user.id == self.admin_id:
logging.debug("Returning admin state: %s", admin_state)
return self.admin_id, admin_state

# If the user state is already non-idle, we are already in the business logic
if (user_state := self._states.get(user.id, State.IDLE)) != State.IDLE:
logging.debug("Returning user state: %s", user_state)
return user.id, user_state

# On first interaction, we need to determine what to do with the user
# if the admin is not idle, we put the user on hold. Otherwise, they may send the first
# message, and we put the admin in waiting for reply to avoid another user occupying the
# admin first
effective_user_state = self.HOLD if admin_state != State.IDLE else self.WELCOMING
self.do_set_state(user.id, effective_user_state)
if effective_user_state == self.WELCOMING:
self.do_set_state(self.admin_id, self.WAITING_FOR_REPLY)

logging.debug("Returning user state: %s", effective_user_state)
return user.id, effective_user_state

def do_set_state(self, key: int, state: State) -> None:
key_name = "admin" if key == self.admin_id else key
logging.debug("Setting %s state to %s", key_name, state)
self._states[key] = state


async def welcome_user(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
await update.effective_message.forward(context.bot_data["admin_id"])
await update.effective_message.reply_text(
"Welcome! Your message has been forwarded to the admin. They will get back to you soon.",
)
await context.set_state(UserSupportMachine.WAITING_FOR_REPLY)
await context.fsm.set_state(context.bot_data["admin_id"], UserSupportMachine.WRITING)
context.bot_data["active_user"] = update.effective_user.id


async def conversation_timeout(context: ContextTypes.DEFAULT_TYPE) -> None:
active_user = context.bot_data.get("active_user")
admin_id = context.bot_data["admin_id"]

async def handle(user_id: int) -> None:
await context.bot.send_message(
user_id, "The conversation has been stopped due to inactivity."
)
await context.fsm.set_state(user_id, State.IDLE)

if active_user:
await handle(active_user)
await handle(admin_id)


async def handle_reply(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
# Cancel the conversation timeout
if job := context.bot_data.get("conversation_timeout"):
job.schedule_removal()

if not (active_user := context.bot_data.get("active_user")):
logger.warning("No active user found, ignoring message")

>
active_user
if update.effective_user.id == (admin_id := context.bot_data["admin_id"])
else admin_id
)
await context.bot.send_message(target, update.effective_message.text)
logging.debug("Forwarded message to %s", target)
await context.set_state(UserSupportMachine.WAITING_FOR_REPLY)
logging.debug("Done setting state to WAITING_FOR_REPLY for %s", target)
await context.fsm.set_state(target, UserSupportMachine.WRITING)
logging.debug("Done setting state to WRITING for %s, context.fsm_key")

# Reset the conversation timeout
job = context.job_queue.run_once(conversation_timeout, 30)
context.bot_data["conversation_timeout"] = job


async def stop_conversation(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
text = "The conversation has been stopped."
admin_id = context.bot_data["admin_id"]
active_user = context.bot_data.get("active_user")

await context.bot.send_message(admin_id, text)
await context.fsm.set_state(admin_id, State.IDLE)
if active_user:
await context.bot.send_message(active_user, text)
await context.fsm.set_state(active_user, State.IDLE)


async def hold_melody(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
await update.effective_message.reply_text(
"You have been put on hold. The admin will get back to you soon. Please hear some music "
"while you wait: https://www.youtube.com/watch?v=dQw4w9WgXcQ"
)


async def not_your_turn(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
await update.effective_message.reply_text(
"It's not your turn yet. Please wait for the other party to reply to your message."
)


def main() -> None:
application = Application.builder().token("TOKEN").build()
application.fsm = UserSupportMachine(admin_id=123456789)
application.bot_data["admin_id"] = application.fsm.admin_id

# Users are welcomed only if they are in the corresponding state
application.add_handler(
MessageHandler(~filters.User(application.fsm.admin_id), welcome_user),
state=UserSupportMachine.WELCOMING,
)

# Conversation logic:
# * forward messages between user and admin
# * stop the conversation at any time (admin or user)
# * point out that the other party is currently writing
application.add_handler(
MessageHandler(filters.ALL, handle_reply), state=UserSupportMachine.WRITING
)
application.add_handler(
CommandHandler("stop", stop_conversation),
state=UserSupportMachine.WAITING_FOR_REPLY | UserSupportMachine.WRITING,
)
application.add_handler(
MessageHandler(filters.ALL, not_your_turn), state=UserSupportMachine.WAITING_FOR_REPLY
)

# If the admin is busy, put the user on hold
application.add_handler(
MessageHandler(filters.ALL, hold_melody), state=UserSupportMachine.HOLD
)

application.run_polling(allowed_updates=Update.ALL_TYPES)


if __name__ == "__main__":
main()
4 changes: 4 additions & 0 deletions telegram/ext/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@
"Defaults",
"DictPersistence",
"ExtBot",
"FiniteStateMachine",
"InlineQueryHandler",
"InvalidCallbackData",
"Job",
Expand All @@ -57,6 +58,8 @@
"PrefixHandler",
"ShippingQueryHandler",
"SimpleUpdateProcessor",
"SingleStateMachine",
"State",
"StringCommandHandler",
"StringRegexHandler",
"TypeHandler",
Expand All @@ -77,6 +80,7 @@
from ._defaults import Defaults
from ._dictpersistence import DictPersistence
from ._extbot import ExtBot
from ._fsm import FiniteStateMachine, SingleStateMachine, State
from ._handlers.basehandler import BaseHandler
from ._handlers.businessconnectionhandler import BusinessConnectionHandler
from ._handlers.businessmessagesdeletedhandler import BusinessMessagesDeletedHandler
Expand Down
44 changes: 32 additions & 12 deletions telegram/ext/_application.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@
from telegram.ext._basepersistence import BasePersistence
from telegram.ext._contexttypes import ContextTypes
from telegram.ext._extbot import ExtBot
from telegram.ext._fsm import SingleStateMachine, State
from telegram.ext._handlers.basehandler import BaseHandler
from telegram.ext._updater import Updater
from telegram.ext._utils.stack import was_called_by
Expand All @@ -59,7 +60,7 @@
from socket import socket

from telegram import Message
from telegram.ext import ConversationHandler, JobQueue
from telegram.ext import ConversationHandler, FiniteStateMachine, JobQueue
from telegram.ext._applicationbuilder import InitApplicationBuilder
from telegram.ext._baseupdateprocessor import BaseUpdateProcessor
from telegram.ext._jobqueue import Job
Expand Down Expand Up @@ -266,6 +267,7 @@ class Application(
"update_queue",
"updater",
"user_data",
"fsm",
)
# Allowing '__weakref__' creation here since we need it for the JobQueue
# Currently the __weakref__ slot is already created
Expand Down Expand Up @@ -301,11 +303,12 @@ def __init__(
stacklevel=2,
)

self.fsm: FiniteStateMachine = SingleStateMachine()
self.bot: BT = bot
self.update_queue: asyncio.Queue[object] = update_queue
self.context_types: ContextTypes[CCT, UD, CD, BD] = context_types
self.updater: Optional[Updater] = updater
self.handlers: dict[int, list[BaseHandler[Any, CCT, Any]]] = {}
self.handlers: dict[State, dict[int, list[BaseHandler[Any, CCT, Any]]]] = {}
self.error_handlers: dict[
HandlerCallback[object, CCT, None], Union[bool, DefaultValue[bool]]
] = {}
Expand Down Expand Up @@ -1280,8 +1283,17 @@ async def process_update(self, update: object) -> None:

context = None
any_blocking = False # Flag which is set to True if any handler specifies block=True
fsm_key, fsm_state = self.fsm.get_active_key_state(update)

for handlers in self.handlers.values():
for state, state_handlers_ in self.handlers.items():
if state.matches(fsm_state):
state_handlers = state_handlers_
break
else:
_LOGGER.debug("No handlers found for key %s in state %s", fsm_key, fsm_state)
return

for handlers in state_handlers.values():
try:
for handler in handlers:
check = handler.check_update(update) # Should the handler handle this update?
Expand All @@ -1291,6 +1303,8 @@ async def process_update(self, update: object) -> None:
if not context: # build a context if not already built
try:
context = self.context_types.context.from_update(update, self)
context.state = fsm_state
context.fsm_key = fsm_key
except Exception as exc:
_LOGGER.critical(
(
Expand Down Expand Up @@ -1340,7 +1354,12 @@ async def process_update(self, update: object) -> None:
# (in __create_task_callback)
self._mark_for_persistence_update(update=update)

def add_handler(self, handler: BaseHandler[Any, CCT, Any], group: int = DEFAULT_GROUP) -> None:
def add_handler(
self,
handler: BaseHandler[Any, CCT, Any],
group: int = DEFAULT_GROUP,
state: State = State.IDLE,
) -> None:
"""Register a handler.

TL;DR: Order and priority counts. 0 or 1 handlers per group will be used. End handling of
Expand Down Expand Up @@ -1399,11 +1418,11 @@ def add_handler(self, handler: BaseHandler[Any, CCT, Any], group: int = DEFAULT_
stacklevel=2,
)

if group not in self.handlers:
self.handlers[group] = []
self.handlers = dict(sorted(self.handlers.items())) # lower -> higher groups
state_handlers = self.handlers.setdefault(state, {})
if group not in state_handlers:
state_handlers[group] = []

self.handlers[group].append(handler)
state_handlers[group].append(handler)

def add_handlers(
self,
Expand Down Expand Up @@ -1475,10 +1494,11 @@ def remove_handler(
group (:obj:`object`, optional): The group identifier. Default is ``0``.

"""
if handler in self.handlers[group]:
self.handlers[group].remove(handler)
if not self.handlers[group]:
del self.handlers[group]
for state_handlers in self.handlers.values():
if handler in state_handlers[group]:
state_handlers[group].remove(handler)
if not state_handlers[group]:
del state_handlers[group]
Comment on lines +1518 to +1519

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Does this remove the group if its empty? Can we comment it

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

yes, that's basically the same as state_handlers.pop(group, None), I think. It will call state_handlers.__del__(group). IMHO this is basic python syntax 😬


def drop_chat_data(self, chat_id: int) -> None:
"""Drops the corresponding entry from the :attr:`chat_data`. Will also be deleted from
Expand Down
Loading