diff --git a/app.py b/app.py index daf465b..9120248 100644 --- a/app.py +++ b/app.py @@ -65,11 +65,19 @@ def on_join(data): for message in previous_messages: emit( "previous_messages", - {"username": message.username, "message": message.content}, + { + "id": message.id, + "username": message.username, + "message": message.content, + }, room=request.sid, ) - emit("message", f"{data['username']} has joined the room.", room=room) + emit( + "message", + {"id": None, "content": f"{data['username']} has joined the room."}, + room=room, + ) @socketio.on("message") @@ -81,30 +89,31 @@ def handle_message(data): db.session.add(new_message) db.session.commit() - emit("message", f"{data['username']}: {data['message']}", room=data["room"]) + emit( + "message", + {"id": new_message.id, "content": f"{data['username']}: {data['message']}"}, + room=data["room"], + ) - if "claude" in data["message"] or "gpt" in data["message"]: + if "claude" in data["message"] or "gpt" in data["message"]: + # Emit a temporary message indicating that llm is processing + emit( + "message", + {"id": None, "content": f"Processing..."}, + room=data["room"], + ) if "claude" in data["message"]: - func = chat_claude - elif "gpt" in data["message"]: - func = chat_gpt - - # Emit a temporary message indicating that llm is processing - emit("message", f"Processing...", room=data["room"]) - - # Call the chat_claude function without blocking using eventlet.spawn - eventlet.spawn(func, data["username"], data["room"], data["message"]) + eventlet.spawn(chat_claude, data["username"], data["room"], data["message"]) + if "gpt" in data["message"]: + eventlet.spawn(chat_gpt, data["username"], data["room"], data["message"]) -def chat_claude(username, room, message): - +def chat_claude(username, room, message): with app.app_context(): # claude has a 100,000 token context window for prompts. all_messages = ( - Message.query.filter_by(room=room) - .order_by(Message.id.desc()) - .all() + Message.query.filter_by(room=room).order_by(Message.id.desc()).all() ) chat_history = "" @@ -114,10 +123,9 @@ def chat_claude(username, room, message): chat_history += f"Assistant: {msg.username}: {msg.content}\n\n" else: chat_history += f"Human: {msg.username}: {msg.content}\n\n" - + # append the new message. chat_history += f"Human: {username}: {message}\n\nAssistant:" - # Initialize the Bedrock client using boto3 client = boto3.client("bedrock-runtime", region_name="us-east-1") @@ -146,9 +154,15 @@ def chat_claude(username, room, message): # Process the event stream buffer = "" + # save empty message, we need the ID when we chunk the response. + with app.app_context(): + new_message = Message(username="anthropic.claude-v2", content=buffer, room=room) + db.session.add(new_message) + db.session.commit() + msg_id = new_message.id + first_chunk = True for event in response["body"]: - content = "" if "chunk" in event: @@ -159,15 +173,27 @@ def chat_claude(username, room, message): buffer += content # Accumulate content if first_chunk: - socketio.emit("message_chunk", f"{username} (anthropic.claude-v2): {content}", room=room) + socketio.emit( + "message_chunk", + { + "id": msg_id, + "content": f"{username} (anthropic.claude-v2): {content}", + }, + room=room, + ) first_chunk = False else: - socketio.emit("message_chunk", content, room=room) + socketio.emit( + "message_chunk", + {"id": msg_id, "content": content}, + room=room, + ) socketio.sleep(0) # Force immediate handling - + # Save the entire completion to the database with app.app_context(): - new_message = Message(username="anthropic.claude-v2", content=buffer, room=room) + new_message = db.session.query(Message).filter(Message.id == msg_id).one_or_none() + new_message.content = buffer db.session.add(new_message) db.session.commit() @@ -175,7 +201,6 @@ def chat_claude(username, room, message): def chat_gpt(username, room, message): - with app.app_context(): last_messages = ( Message.query.filter_by(room=room) @@ -185,7 +210,14 @@ def chat_gpt(username, room, message): ) chat_history = [ - {"role": "system" if (msg.username == "gpt-3.5-turbo" or msg.username == "anthropic.claude-v2") else "user", "content": f"{msg.username}: {msg.content}"} + { + "role": "system" + if ( + msg.username == "gpt-3.5-turbo" or msg.username == "anthropic.claude-v2" + ) + else "user", + "content": f"{msg.username}: {msg.content}", + } for msg in reversed(last_messages) ] @@ -193,6 +225,13 @@ def chat_gpt(username, room, message): buffer = "" # Content buffer for accumulating the chunks + # save empty message, we need the ID when we chunk the response. + with app.app_context(): + new_message = Message(username="gpt-3.5-turbo", content=buffer, room=room) + db.session.add(new_message) + db.session.commit() + msg_id = new_message.id + first_chunk = True for chunk in openai.ChatCompletion.create( model="gpt-3.5-turbo", @@ -206,15 +245,27 @@ def chat_gpt(username, room, message): buffer += content # Accumulate content if first_chunk: - socketio.emit("message_chunk", f"{username} (gpt-3.5-turbo): {content}", room=room) + socketio.emit( + "message_chunk", + { + "id": msg_id, + "content": f"{username} (gpt-3.5-turbo): {content}", + }, + room=room, + ) first_chunk = False else: - socketio.emit("message_chunk", content, room=room) + socketio.emit( + "message_chunk", + {"id": msg_id, "content": content}, + room=room, + ) socketio.sleep(0) # Force immediate handling # Save the entire completion to the database with app.app_context(): - new_message = Message(username="gpt-3.5-turbo", content=buffer, room=room) + new_message = db.session.query(Message).filter(Message.id == msg_id).one_or_none() + new_message.content = buffer db.session.add(new_message) db.session.commit() diff --git a/templates/chat.html b/templates/chat.html index 4962560..e8fb1c4 100644 --- a/templates/chat.html +++ b/templates/chat.html @@ -70,9 +70,6 @@ const urlParams = new URLSearchParams(window.location.search); const username = urlParams.get('username'); const room = '{{ room }}'; -let lastMessageElement = null; -let lastMessageContent = ""; - // Function to handle sending the message function sendMessage() { const message = document.getElementById('message').value; @@ -100,63 +97,92 @@ socket.on('connect', () => { socket.emit('join', {username: username, room: room}); }); -socket.on('message', (message) => { +socket.on('message', (data) => { const newMessage = document.createElement('p'); - newMessage.innerHTML = marked.marked(message); + + // Set the message ID as the id attribute + newMessage.id = "message-" + data.id; // Prefixing with "message-" to ensure the ID starts with a letter + + // Convert the message content to markdown and set it as the innerHTML + newMessage.innerHTML = marked.marked(data.content); + + // Append the new message to the chat container document.getElementById('chat').appendChild(newMessage); + // Apply syntax highlighting to any code blocks within the message newMessage.querySelectorAll('pre code').forEach((block) => { hljs.highlightElement(block); }); + // Scroll to the bottom of the chat container document.getElementById('chat').scrollTop = document.getElementById('chat').scrollHeight; }); + socket.on('previous_messages', (data) => { const chat = document.getElementById('chat'); - chat.innerHTML += '
' + data.username + ': ' + marked.marked(data.message) + '
'; - chat.querySelectorAll('pre code').forEach((block) => { - hljs.highlightBlock(block); + + // Create a new message element + const newMessage = document.createElement('p'); + + // Set the ID attribute for the message + newMessage.id = "message-" + data.id; + + // Set the inner content for the message + newMessage.innerHTML = `${data.username}: ${marked.marked(data.message)}`; + + // Append the new message to the chat container + chat.appendChild(newMessage); + + // Apply syntax highlighting to any code blocks within the message + newMessage.querySelectorAll('pre code').forEach((block) => { + hljs.highlightElement(block); }); }); + socket.on('delete_processing_message', (data) => { const tempMessage = document.getElementById("processing"); if (tempMessage) { tempMessage.remove(); } - lastMessageElement = null; - lastMessageContent = ""; }); -socket.on('message_chunk', (message_chunk) => { - if (lastMessageElement) { - // Accumulate the new chunk - lastMessageContent += message_chunk; - // Render the entire accumulated content as Markdown - lastMessageElement.innerHTML = marked.marked(lastMessageContent); +// A dictionary to hold buffers for each message ID +const messageBuffers = {}; - // Apply syntax highlighting to code blocks within the content - lastMessageElement.querySelectorAll('pre code').forEach((block) => { - hljs.highlightElement(block); - }); - } else { - lastMessageContent = message_chunk; - const newMessage = document.createElement('p'); - newMessage.innerHTML = marked.marked(lastMessageContent); - document.getElementById('chat').appendChild(newMessage); - lastMessageElement = newMessage; +socket.on('message_chunk', (data) => { + const messageId = "message-" + data.id; + let targetMessageElement = document.getElementById(messageId); - // Apply syntax highlighting to code blocks within the content - newMessage.querySelectorAll('pre code').forEach((block) => { - hljs.highlightElement(block); - }); + // If the target message element doesn't exist, it's an initial chunk + if (!targetMessageElement) { + targetMessageElement = document.createElement('p'); + targetMessageElement.id = messageId; + document.getElementById('chat').appendChild(targetMessageElement); } + // If the message buffer for this ID doesn't exist, create it + if (!messageBuffers[data.id]) { + messageBuffers[data.id] = ""; + } + + // Append the chunk to the buffer + messageBuffers[data.id] += data.content; + + // Process the entire buffer with marked and set it as the content of the target element + targetMessageElement.innerHTML = marked.marked(messageBuffers[data.id]); + + // Apply syntax highlighting to code blocks within the content + targetMessageElement.querySelectorAll('pre code').forEach((block) => { + hljs.highlightElement(block); + }); + // Scroll to the bottom of the chat container document.getElementById('chat').scrollTop = document.getElementById('chat').scrollHeight; }); +