From 8beacd5ac8473105a24717950f7595002b8bea16 Mon Sep 17 00:00:00 2001 From: Russell Ballestrini Date: Sun, 3 Dec 2023 12:47:46 -0500 Subject: [PATCH] new /cancel command to stop generation in the middle of streaming to chat modified: README.rst modified: app.py --- README.rst | 1 + app.py | 45 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 46 insertions(+) diff --git a/README.rst b/README.rst index 2a681de..dba6485 100644 --- a/README.rst +++ b/README.rst @@ -102,6 +102,7 @@ The application supports special commands for interacting with AWS S3: - ``/s3 load ``: Loads a file from S3 and displays its content in the chatroom. - ``/s3 save ``: Saves the most recent code block from the chatroom to S3. - ``/title new``: Generates a new title which reflects conversation content for the current chatroom using gpt-4. +- ``/cancel``: cancel the most recent chat completion from streaming into chatroom. Contributing ------------ diff --git a/app.py b/app.py index 0b6ab63..e9db0e0 100644 --- a/app.py +++ b/app.py @@ -36,6 +36,8 @@ db = SQLAlchemy(app) # profile_name = args.profile profile_name = None +# Global dictionary to keep track of cancellation requests +cancellation_requests = {} class Room(db.Model): id = db.Column(db.Integer, primary_key=True) @@ -293,6 +295,31 @@ def load_s3_file(room_name, s3_file_path, username): room=room_name, ) +def cancel_generation(room_name, username): + with app.app_context(): + room = get_room(room_name) + # Get the most recent message for the room that is being generated + latest_message = ( + Message.query.filter_by(room_id=room.id) + .order_by(Message.id.desc()) + .offset(1) + .first() + ) + + if latest_message: + # Set the cancellation request for the given message ID + cancellation_requests[latest_message.id] = True + # Optionally, inform the user that the generation has been canceled + socketio.emit( + "message", + { + "id": None, + "username": "System", + "content": f"Generation for message ID {latest_message.id} has been canceled.", + }, + room=room_name, + ) + @socketio.on("message") def handle_message(data): @@ -336,6 +363,9 @@ def handle_message(data): ) if command.startswith("/title new"): eventlet.spawn(generate_new_title, room_name, data["username"]) + if command.startswith("/cancel"): + # Cancel the most recent generation request + eventlet.spawn(cancel_generation, room_name, data["username"]) if ( "claude-v1" in data["message"] @@ -448,10 +478,17 @@ def chat_claude(username, room_name, message, model_name="anthropic.claude-v1"): db.session.commit() msg_id = new_message.id + cancellation_requests[msg_id] = False + first_chunk = True for event in response["body"]: content = "" + # Check if there has been a cancellation request, break if there is. + if cancellation_requests.get(msg_id): + del cancellation_requests[msg_id] + break + if "chunk" in event: chunk_data = json.loads(event["chunk"]["bytes"].decode()) content = chunk_data["completion"] @@ -540,10 +577,18 @@ def chat_gpt(username, room_name, message, model_name="gpt-3.5-turbo"): db.session.commit() msg_id = new_message.id + cancellation_requests[msg_id] = False + first_chunk = True for chunk in openai_client.chat.completions.create( model=model_name, messages=chat_history, temperature=0, stream=True ): + + # Check if there has been a cancellation request, break if there is. + if cancellation_requests.get(msg_id): + del cancellation_requests[msg_id] + break + content = chunk.choices[0].delta.content if content: