Coverage for open_webui/models/messages.py: 32%
276 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1import time
2import uuid
3from typing import Optional
5from open_webui.internal.db import Base, JSONField, get_async_db_context
6from open_webui.models.channels import ChannelMember, Channels
7from open_webui.models.tags import Tag, TagModel, Tags
8from open_webui.models.users import User, UserNameResponse, Users
9from pydantic import BaseModel, ConfigDict, field_validator
10from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, and_, delete, func, or_, select, text
11from sqlalchemy.ext.asyncio import AsyncSession
12from sqlalchemy.sql import exists
14####################
15# Message DB Schema
16####################
19class MessageReaction(Base):
20 __tablename__ = 'message_reaction'
21 id = Column(Text, primary_key=True, unique=True)
22 user_id = Column(Text)
23 message_id = Column(Text)
24 name = Column(Text)
25 created_at = Column(BigInteger)
28class MessageReactionModel(BaseModel):
29 model_config = ConfigDict(from_attributes=True)
31 id: str
32 user_id: str
33 message_id: str
34 name: str
35 created_at: int # timestamp in epoch
38class Message(Base):
39 __tablename__ = 'message'
40 id = Column(Text, primary_key=True, unique=True)
42 user_id = Column(Text)
43 channel_id = Column(Text, nullable=True)
45 reply_to_id = Column(Text, nullable=True)
46 parent_id = Column(Text, nullable=True)
48 # Pins
49 is_pinned = Column(Boolean, nullable=False, default=False)
50 pinned_at = Column(BigInteger, nullable=True)
51 pinned_by = Column(Text, nullable=True)
53 content = Column(Text)
54 data = Column(JSON, nullable=True)
55 meta = Column(JSON, nullable=True)
57 created_at = Column(BigInteger) # time_ns
58 updated_at = Column(BigInteger) # time_ns
61class MessageModel(BaseModel):
62 model_config = ConfigDict(from_attributes=True)
64 id: str
65 user_id: str
66 channel_id: Optional[str] = None
68 reply_to_id: Optional[str] = None
69 parent_id: Optional[str] = None
71 # Pins
72 is_pinned: bool = False
73 pinned_by: Optional[str] = None
74 pinned_at: Optional[int] = None # timestamp in epoch (time_ns)
76 content: str
77 data: Optional[dict] = None
78 meta: Optional[dict] = None
80 created_at: int # timestamp in epoch (time_ns)
81 updated_at: int # timestamp in epoch (time_ns)
84####################
85# Forms
86####################
89class MessageForm(BaseModel):
90 temp_id: Optional[str] = None
91 content: str
92 reply_to_id: Optional[str] = None
93 parent_id: Optional[str] = None
94 data: Optional[dict] = None
95 meta: Optional[dict] = None
98class Reactions(BaseModel):
99 name: str
100 users: list[dict]
101 count: int
104class MessageUserResponse(MessageModel):
105 user: Optional[UserNameResponse] = None
108class MessageUserSlimResponse(MessageUserResponse):
109 data: bool | None = None
111 @field_validator('data', mode='before')
112 def convert_data_to_bool(cls, v):
113 # No data or not a dict → False
114 if not isinstance(v, dict):
115 return False
117 # True if ANY value in the dict is non-empty
118 return any(bool(val) for val in v.values())
121class MessageReplyToResponse(MessageUserResponse):
122 reply_to_message: Optional[MessageUserSlimResponse] = None
125class MessageWithReactionsResponse(MessageUserSlimResponse):
126 reactions: list[Reactions]
129class MessageResponse(MessageReplyToResponse):
130 latest_reply_at: Optional[int]
131 reply_count: int
132 reactions: list[Reactions]
135class MessageTable:
136 async def insert_new_message(
137 self,
138 form_data: MessageForm,
139 channel_id: str,
140 user_id: str,
141 db: Optional[AsyncSession] = None,
142 ) -> Optional[MessageModel]:
143 async with get_async_db_context(db) as db:
144 channel_member = await Channels.join_channel(channel_id, user_id)
146 id = str(uuid.uuid4())
147 ts = int(time.time_ns())
149 message = MessageModel(
150 **{
151 'id': id,
152 'user_id': user_id,
153 'channel_id': channel_id,
154 'reply_to_id': form_data.reply_to_id,
155 'parent_id': form_data.parent_id,
156 'is_pinned': False,
157 'pinned_at': None,
158 'pinned_by': None,
159 'content': form_data.content,
160 'data': form_data.data,
161 'meta': form_data.meta,
162 'created_at': ts,
163 'updated_at': ts,
164 }
165 )
166 result = Message(**message.model_dump())
168 db.add(result)
169 await db.commit()
170 await db.refresh(result)
171 return MessageModel.model_validate(result) if result else None
173 async def get_message_by_id(
174 self,
175 id: str,
176 include_thread_replies: Optional[bool] = True,
177 db: Optional[AsyncSession] = None,
178 ) -> Optional[MessageResponse]:
179 async with get_async_db_context(db) as db:
180 message = await db.get(Message, id)
181 if not message:
182 return None
184 reply_to_message = (
185 await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
186 if message.reply_to_id
187 else None
188 )
190 reactions = await self.get_reactions_by_message_id(id, db=db)
192 thread_replies = []
193 if include_thread_replies:
194 thread_replies = await self.get_thread_replies_by_message_id(id, db=db)
196 # Check if message was sent by webhook (webhook info in meta takes precedence)
197 webhook_info = message.meta.get('webhook') if message.meta else None
198 if webhook_info and webhook_info.get('id'):
199 # Look up webhook by ID to get current name
200 webhook = await Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
201 if webhook:
202 user_info = {
203 'id': webhook.id,
204 'name': webhook.name,
205 'role': 'webhook',
206 }
207 else:
208 # Webhook was deleted, use placeholder
209 user_info = {
210 'id': webhook_info.get('id'),
211 'name': 'Deleted Webhook',
212 'role': 'webhook',
213 }
214 else:
215 user = await Users.get_user_by_id(message.user_id, db=db)
216 user_info = user.model_dump() if user else None
218 return MessageResponse.model_validate(
219 {
220 **MessageModel.model_validate(message).model_dump(),
221 'user': user_info,
222 'reply_to_message': (reply_to_message.model_dump() if reply_to_message else None),
223 'latest_reply_at': (thread_replies[0].created_at if thread_replies else None),
224 'reply_count': len(thread_replies),
225 'reactions': reactions,
226 }
227 )
229 async def _resolve_user_info(self, message: Message, db: AsyncSession) -> Optional[dict]:
230 """Resolve user info from message, handling webhook messages."""
231 webhook_info = message.meta.get('webhook') if message.meta else None
232 if webhook_info and webhook_info.get('id'):
233 webhook = await Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
234 if webhook:
235 return {
236 'id': webhook.id,
237 'name': webhook.name,
238 'role': 'webhook',
239 }
240 else:
241 return {
242 'id': webhook_info.get('id'),
243 'name': 'Deleted Webhook',
244 'role': 'webhook',
245 }
246 return None
248 async def get_thread_replies_by_message_id(
249 self, id: str, db: Optional[AsyncSession] = None
250 ) -> list[MessageReplyToResponse]:
251 async with get_async_db_context(db) as db:
252 result = await db.execute(select(Message).filter_by(parent_id=id).order_by(Message.created_at.desc()))
253 all_messages = result.scalars().all()
255 messages = []
256 for message in all_messages:
257 reply_to_message = (
258 await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
259 if message.reply_to_id
260 else None
261 )
263 user_info = await self._resolve_user_info(message, db)
265 messages.append(
266 MessageReplyToResponse.model_validate(
267 {
268 **MessageModel.model_validate(message).model_dump(),
269 'user': user_info,
270 'reply_to_message': (reply_to_message.model_dump() if reply_to_message else None),
271 }
272 )
273 )
274 return messages
276 async def get_reply_user_ids_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[str]:
277 async with get_async_db_context(db) as db:
278 result = await db.execute(select(Message.user_id).filter_by(parent_id=id))
279 return [row[0] for row in result.all()]
281 async def get_messages_by_channel_id(
282 self,
283 channel_id: str,
284 skip: int = 0,
285 limit: int = 50,
286 db: Optional[AsyncSession] = None,
287 ) -> list[MessageReplyToResponse]:
288 async with get_async_db_context(db) as db:
289 result = await db.execute(
290 select(Message)
291 .filter_by(channel_id=channel_id, parent_id=None)
292 .order_by(Message.created_at.desc())
293 .offset(skip)
294 .limit(limit)
295 )
296 all_messages = result.scalars().all()
298 messages = []
299 for message in all_messages:
300 reply_to_message = (
301 await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
302 if message.reply_to_id
303 else None
304 )
306 user_info = await self._resolve_user_info(message, db)
308 messages.append(
309 MessageReplyToResponse.model_validate(
310 {
311 **MessageModel.model_validate(message).model_dump(),
312 'user': user_info,
313 'reply_to_message': (reply_to_message.model_dump() if reply_to_message else None),
314 }
315 )
316 )
317 return messages
319 async def get_messages_by_parent_id(
320 self,
321 channel_id: str,
322 parent_id: str,
323 skip: int = 0,
324 limit: int = 50,
325 db: Optional[AsyncSession] = None,
326 ) -> list[MessageReplyToResponse]:
327 async with get_async_db_context(db) as db:
328 message = await db.get(Message, parent_id)
330 # Thread parent must belong to the requested channel; never disclose a foreign-channel message.
331 if not message or message.channel_id != channel_id:
332 return []
334 result = await db.execute(
335 select(Message)
336 .filter_by(channel_id=channel_id, parent_id=parent_id)
337 .order_by(Message.created_at.desc())
338 .offset(skip)
339 .limit(limit)
340 )
341 all_messages = list(result.scalars().all())
343 # If length of all_messages is less than limit, then add the parent message
344 if len(all_messages) < limit:
345 all_messages.append(message)
347 messages = []
348 for message in all_messages:
349 reply_to_message = (
350 await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
351 if message.reply_to_id
352 else None
353 )
355 user_info = await self._resolve_user_info(message, db)
357 messages.append(
358 MessageReplyToResponse.model_validate(
359 {
360 **MessageModel.model_validate(message).model_dump(),
361 'user': user_info,
362 'reply_to_message': (reply_to_message.model_dump() if reply_to_message else None),
363 }
364 )
365 )
366 return messages
368 async def get_last_message_by_channel_id(
369 self, channel_id: str, db: Optional[AsyncSession] = None
370 ) -> Optional[MessageModel]:
371 async with get_async_db_context(db) as db:
372 result = await db.execute(
373 select(Message).filter_by(channel_id=channel_id).order_by(Message.created_at.desc()).limit(1)
374 )
375 message = result.scalars().first()
376 return MessageModel.model_validate(message) if message else None
378 async def get_pinned_messages_by_channel_id(
379 self,
380 channel_id: str,
381 skip: int = 0,
382 limit: int = 50,
383 db: Optional[AsyncSession] = None,
384 ) -> list[MessageModel]:
385 async with get_async_db_context(db) as db:
386 result = await db.execute(
387 select(Message)
388 .filter_by(channel_id=channel_id, is_pinned=True)
389 .order_by(Message.pinned_at.desc())
390 .offset(skip)
391 .limit(limit)
392 )
393 all_messages = result.scalars().all()
394 return [MessageModel.model_validate(message) for message in all_messages]
396 async def update_message_by_id(
397 self, id: str, form_data: MessageForm, db: Optional[AsyncSession] = None
398 ) -> Optional[MessageModel]:
399 async with get_async_db_context(db) as db:
400 message = await db.get(Message, id)
401 message.content = form_data.content
402 message.data = {
403 **(message.data if message.data else {}),
404 **(form_data.data if form_data.data else {}),
405 }
406 message.meta = {
407 **(message.meta if message.meta else {}),
408 **(form_data.meta if form_data.meta else {}),
409 }
410 message.updated_at = int(time.time_ns())
411 await db.commit()
412 await db.refresh(message)
413 return MessageModel.model_validate(message) if message else None
415 async def update_is_pinned_by_id(
416 self,
417 id: str,
418 is_pinned: bool,
419 pinned_by: Optional[str] = None,
420 db: Optional[AsyncSession] = None,
421 ) -> Optional[MessageModel]:
422 async with get_async_db_context(db) as db:
423 message = await db.get(Message, id)
424 message.is_pinned = is_pinned
425 message.pinned_at = int(time.time_ns()) if is_pinned else None
426 message.pinned_by = pinned_by if is_pinned else None
427 await db.commit()
428 await db.refresh(message)
429 return MessageModel.model_validate(message) if message else None
431 async def get_unread_message_count(
432 self,
433 channel_id: str,
434 user_id: str,
435 last_read_at: Optional[int] = None,
436 db: Optional[AsyncSession] = None,
437 ) -> int:
438 async with get_async_db_context(db) as db:
439 stmt = select(func.count(Message.id)).filter(
440 Message.channel_id == channel_id,
441 Message.parent_id == None, # only count top-level messages
442 Message.created_at > (last_read_at if last_read_at else 0),
443 )
444 if user_id:
445 stmt = stmt.filter(Message.user_id != user_id)
446 result = await db.execute(stmt)
447 return result.scalar()
449 async def add_reaction_to_message(
450 self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None
451 ) -> Optional[MessageReactionModel]:
452 async with get_async_db_context(db) as db:
453 # check for existing reaction
454 result = await db.execute(select(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name))
455 existing_reaction = result.scalars().first()
456 if existing_reaction:
457 return MessageReactionModel.model_validate(existing_reaction)
459 reaction_id = str(uuid.uuid4())
460 reaction = MessageReactionModel(
461 id=reaction_id,
462 user_id=user_id,
463 message_id=id,
464 name=name,
465 created_at=int(time.time_ns()),
466 )
467 result = MessageReaction(**reaction.model_dump())
468 db.add(result)
469 await db.commit()
470 await db.refresh(result)
471 return MessageReactionModel.model_validate(result) if result else None
473 async def get_reactions_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[Reactions]:
474 async with get_async_db_context(db) as db:
475 # JOIN User so all user info is fetched in one query
476 result = await db.execute(
477 select(MessageReaction, User)
478 .join(User, MessageReaction.user_id == User.id)
479 .filter(MessageReaction.message_id == id)
480 )
481 results = result.all()
483 reactions = {}
485 for reaction, user in results:
486 if reaction.name not in reactions:
487 reactions[reaction.name] = {
488 'name': reaction.name,
489 'users': [],
490 'count': 0,
491 }
493 reactions[reaction.name]['users'].append(
494 {
495 'id': user.id,
496 'name': user.name,
497 }
498 )
499 reactions[reaction.name]['count'] += 1
501 return [Reactions(**reaction) for reaction in reactions.values()]
503 async def get_reactions_by_message_ids(
504 self, ids: list[str], db: Optional[AsyncSession] = None
505 ) -> dict[str, list[Reactions]]:
506 """Batch-fetch reactions for multiple messages in a single query.
508 Returns a dict mapping each message_id to its list of Reactions.
509 Messages with no reactions map to an empty list.
510 """
511 if not ids:
512 return {}
514 async with get_async_db_context(db) as db:
515 result = await db.execute(
516 select(MessageReaction, User)
517 .join(User, MessageReaction.user_id == User.id)
518 .filter(MessageReaction.message_id.in_(ids))
519 )
520 rows = result.all()
522 # Group by (message_id, reaction_name)
523 grouped: dict[str, dict[str, dict]] = {mid: {} for mid in ids}
524 for reaction, user in rows:
525 mid = reaction.message_id
526 if mid not in grouped:
527 grouped[mid] = {}
528 if reaction.name not in grouped[mid]:
529 grouped[mid][reaction.name] = {
530 'name': reaction.name,
531 'users': [],
532 'count': 0,
533 }
534 grouped[mid][reaction.name]['users'].append(
535 {
536 'id': user.id,
537 'name': user.name,
538 }
539 )
540 grouped[mid][reaction.name]['count'] += 1
542 return {mid: [Reactions(**r) for r in reactions.values()] for mid, reactions in grouped.items()}
544 async def get_thread_reply_counts_by_message_ids(
545 self, ids: list[str], db: Optional[AsyncSession] = None
546 ) -> dict[str, tuple[int, int | None]]:
547 """Batch-fetch reply counts and latest reply timestamps for multiple parent messages.
549 Returns a dict mapping each parent message_id to a
550 (reply_count, latest_reply_created_at) tuple.
551 Messages with no replies are omitted from the result.
552 """
553 if not ids:
554 return {}
556 async with get_async_db_context(db) as db:
557 result = await db.execute(
558 select(
559 Message.parent_id,
560 func.count(Message.id),
561 func.max(Message.created_at),
562 )
563 .filter(Message.parent_id.in_(ids))
564 .group_by(Message.parent_id)
565 )
566 return {row[0]: (row[1], row[2]) for row in result.all()}
568 async def remove_reaction_by_id_and_user_id_and_name(
569 self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None
570 ) -> bool:
571 async with get_async_db_context(db) as db:
572 await db.execute(delete(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name))
573 await db.commit()
574 return True
576 async def delete_reactions_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
577 async with get_async_db_context(db) as db:
578 await db.execute(delete(MessageReaction).filter_by(message_id=id))
579 await db.commit()
580 return True
582 async def delete_replies_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
583 async with get_async_db_context(db) as db:
584 await db.execute(delete(Message).filter_by(parent_id=id))
585 await db.commit()
586 return True
588 async def delete_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
589 async with get_async_db_context(db) as db:
590 await db.execute(delete(Message).filter_by(id=id))
592 # Delete all reactions to this message
593 await db.execute(delete(MessageReaction).filter_by(message_id=id))
595 await db.commit()
596 return True
598 async def search_messages_by_channel_ids(
599 self,
600 channel_ids: list[str],
601 query: str,
602 start_timestamp: Optional[int] = None,
603 end_timestamp: Optional[int] = None,
604 limit: int = 10,
605 db: Optional[AsyncSession] = None,
606 ) -> list[MessageModel]:
607 """Search messages in specified channels by content."""
608 async with get_async_db_context(db) as db:
609 stmt = select(Message).filter(
610 Message.channel_id.in_(channel_ids),
611 Message.content.ilike(f'%{query}%'),
612 )
614 if start_timestamp:
615 stmt = stmt.filter(Message.created_at >= start_timestamp)
616 if end_timestamp:
617 stmt = stmt.filter(Message.created_at <= end_timestamp)
619 stmt = stmt.order_by(Message.created_at.desc()).limit(limit)
620 result = await db.execute(stmt)
621 messages = result.scalars().all()
622 return [MessageModel.model_validate(msg) for msg in messages]
625Messages = MessageTable()