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

1import time 

2import uuid 

3from typing import Optional 

4 

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 

13 

14#################### 

15# Message DB Schema 

16#################### 

17 

18 

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) 

26 

27 

28class MessageReactionModel(BaseModel): 

29 model_config = ConfigDict(from_attributes=True) 

30 

31 id: str 

32 user_id: str 

33 message_id: str 

34 name: str 

35 created_at: int # timestamp in epoch 

36 

37 

38class Message(Base): 

39 __tablename__ = 'message' 

40 id = Column(Text, primary_key=True, unique=True) 

41 

42 user_id = Column(Text) 

43 channel_id = Column(Text, nullable=True) 

44 

45 reply_to_id = Column(Text, nullable=True) 

46 parent_id = Column(Text, nullable=True) 

47 

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) 

52 

53 content = Column(Text) 

54 data = Column(JSON, nullable=True) 

55 meta = Column(JSON, nullable=True) 

56 

57 created_at = Column(BigInteger) # time_ns 

58 updated_at = Column(BigInteger) # time_ns 

59 

60 

61class MessageModel(BaseModel): 

62 model_config = ConfigDict(from_attributes=True) 

63 

64 id: str 

65 user_id: str 

66 channel_id: Optional[str] = None 

67 

68 reply_to_id: Optional[str] = None 

69 parent_id: Optional[str] = None 

70 

71 # Pins 

72 is_pinned: bool = False 

73 pinned_by: Optional[str] = None 

74 pinned_at: Optional[int] = None # timestamp in epoch (time_ns) 

75 

76 content: str 

77 data: Optional[dict] = None 

78 meta: Optional[dict] = None 

79 

80 created_at: int # timestamp in epoch (time_ns) 

81 updated_at: int # timestamp in epoch (time_ns) 

82 

83 

84#################### 

85# Forms 

86#################### 

87 

88 

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 

96 

97 

98class Reactions(BaseModel): 

99 name: str 

100 users: list[dict] 

101 count: int 

102 

103 

104class MessageUserResponse(MessageModel): 

105 user: Optional[UserNameResponse] = None 

106 

107 

108class MessageUserSlimResponse(MessageUserResponse): 

109 data: bool | None = None 

110 

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 

116 

117 # True if ANY value in the dict is non-empty 

118 return any(bool(val) for val in v.values()) 

119 

120 

121class MessageReplyToResponse(MessageUserResponse): 

122 reply_to_message: Optional[MessageUserSlimResponse] = None 

123 

124 

125class MessageWithReactionsResponse(MessageUserSlimResponse): 

126 reactions: list[Reactions] 

127 

128 

129class MessageResponse(MessageReplyToResponse): 

130 latest_reply_at: Optional[int] 

131 reply_count: int 

132 reactions: list[Reactions] 

133 

134 

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) 

145 

146 id = str(uuid.uuid4()) 

147 ts = int(time.time_ns()) 

148 

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()) 

167 

168 db.add(result) 

169 await db.commit() 

170 await db.refresh(result) 

171 return MessageModel.model_validate(result) if result else None 

172 

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 

183 

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 ) 

189 

190 reactions = await self.get_reactions_by_message_id(id, db=db) 

191 

192 thread_replies = [] 

193 if include_thread_replies: 

194 thread_replies = await self.get_thread_replies_by_message_id(id, db=db) 

195 

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 

217 

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 ) 

228 

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 

247 

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() 

254 

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 ) 

262 

263 user_info = await self._resolve_user_info(message, db) 

264 

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 

275 

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()] 

280 

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() 

297 

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 ) 

305 

306 user_info = await self._resolve_user_info(message, db) 

307 

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 

318 

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) 

329 

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 [] 

333 

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()) 

342 

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) 

346 

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 ) 

354 

355 user_info = await self._resolve_user_info(message, db) 

356 

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 

367 

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 

377 

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] 

395 

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 

414 

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 

430 

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() 

448 

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) 

458 

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 

472 

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() 

482 

483 reactions = {} 

484 

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 } 

492 

493 reactions[reaction.name]['users'].append( 

494 { 

495 'id': user.id, 

496 'name': user.name, 

497 } 

498 ) 

499 reactions[reaction.name]['count'] += 1 

500 

501 return [Reactions(**reaction) for reaction in reactions.values()] 

502 

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. 

507 

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 {} 

513 

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() 

521 

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 

541 

542 return {mid: [Reactions(**r) for r in reactions.values()] for mid, reactions in grouped.items()} 

543 

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. 

548 

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 {} 

555 

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()} 

567 

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 

575 

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 

581 

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 

587 

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)) 

591 

592 # Delete all reactions to this message 

593 await db.execute(delete(MessageReaction).filter_by(message_id=id)) 

594 

595 await db.commit() 

596 return True 

597 

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 ) 

613 

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) 

618 

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] 

623 

624 

625Messages = MessageTable()