make count_tokens public on Message

modified:   app.py
	modified:   migrations/versions/190d5ef26e20_add_token_count_to_message.py
This commit is contained in:
Russell Ballestrini 2023-12-07 10:16:01 -05:00
parent caf4c76ceb
commit c2aaac9768
2 changed files with 9 additions and 10 deletions

14
app.py
View file

@ -48,15 +48,13 @@ class Message(db.Model):
self.username = username
self.content = content
self.room_id = room_id
self.token_count = self._count_tokens()
self.token_count = self.count_tokens()
def _count_tokens(self):
def count_tokens(self):
# Replace 'gpt-3.5-turbo' with the model you are using.
encoding = tiktoken.encoding_for_model("gpt-3.5-turbo")
return len(encoding.encode(self.content))
def _count_tokens(self):
return len(self.content.split())
self.token_count = len(encoding.encode(self.content))
return self.token_count
def is_base64_image(self):
return '<img src="data:image/jpeg;base64,' in self.content
@ -384,6 +382,7 @@ def chat_claude(username, room_name, message, model_name="anthropic.claude-v1"):
)
if new_message:
new_message.content = message_content
new_message.count_tokens()
db.session.add(new_message)
db.session.commit()
socketio.emit(
@ -406,6 +405,7 @@ def chat_claude(username, room_name, message, model_name="anthropic.claude-v1"):
)
if new_message:
new_message.content = buffer
new_message.count_tokens()
db.session.add(new_message)
db.session.commit()
@ -479,6 +479,7 @@ def chat_gpt(username, room_name, message, model_name="gpt-3.5-turbo"):
)
if new_message:
new_message.content = message_content
new_message.count_tokens()
db.session.add(new_message)
db.session.commit()
socketio.emit(
@ -530,6 +531,7 @@ def chat_gpt(username, room_name, message, model_name="gpt-3.5-turbo"):
)
if new_message:
new_message.content = buffer
new_message.count_tokens()
db.session.add(new_message)
db.session.commit()

View file

@ -30,16 +30,13 @@ def upgrade():
# Assuming the content column is not nullable and always has a value
messages = session.query(Message).all()
for message in messages:
message.token_count = len(message.content.split())
message.count_tokens()
session.add(message)
session.commit()
# ### end Alembic commands ###
def downgrade():
# ### commands auto generated by Alembic - please adjust! ###
with op.batch_alter_table('message', schema=None) as batch_op:
batch_op.drop_column('token_count')
# ### end Alembic commands ###