Coverage for open_webui/models/shared_chats.py: 62%

126 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1import logging 

2import time 

3import uuid 

4from typing import Optional 

5 

6from open_webui.internal.db import Base, JSONField, get_async_db_context 

7from pydantic import BaseModel, ConfigDict 

8from sqlalchemy import JSON, BigInteger, Column, ForeignKey, Text, delete, select 

9from sqlalchemy.ext.asyncio import AsyncSession 

10 

11log = logging.getLogger(__name__) 

12 

13#################### 

14# SharedChat DB Schema 

15#################### 

16 

17 

18class SharedChat(Base): 

19 __tablename__ = 'shared_chat' 

20 

21 id = Column(Text, primary_key=True) # The share token (UUID) — used in /s/{id} URL 

22 chat_id = Column(Text, ForeignKey('chat.id', ondelete='CASCADE'), nullable=False) 

23 user_id = Column(Text, nullable=False) # Who created this share 

24 

25 title = Column(Text) 

26 chat = Column(JSON) # Snapshot of chat JSON at share time 

27 

28 created_at = Column(BigInteger) 

29 updated_at = Column(BigInteger) 

30 

31 

32class SharedChatModel(BaseModel): 

33 model_config = ConfigDict(from_attributes=True) 

34 

35 id: str 

36 chat_id: str 

37 user_id: str 

38 

39 title: str 

40 chat: dict 

41 

42 created_at: int 

43 updated_at: int 

44 

45 

46class SharedChatResponse(BaseModel): 

47 id: str 

48 chat_id: str 

49 title: str 

50 share_id: Optional[str] = None # Alias for id, for backward compat 

51 updated_at: int 

52 created_at: int 

53 

54 

55#################### 

56# Table Operations 

57#################### 

58 

59 

60class SharedChatsTable: 

61 async def create(self, chat_id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[SharedChatModel]: 

62 """ 

63 Create a snapshot of the chat for link sharing. 

64 Returns the SharedChatModel with the share token as its id. 

65 """ 

66 async with get_async_db_context(db) as db: 

67 from open_webui.models.chats import Chat 

68 

69 chat = await db.get(Chat, chat_id) 

70 if not chat: 

71 return None 

72 

73 share_id = str(uuid.uuid4()) 

74 now = int(time.time()) 

75 

76 shared_chat = SharedChat( 

77 id=share_id, 

78 chat_id=chat_id, 

79 user_id=user_id, 

80 title=chat.title, 

81 chat=chat.chat, 

82 created_at=now, 

83 updated_at=now, 

84 ) 

85 db.add(shared_chat) 

86 await db.commit() 

87 await db.refresh(shared_chat) 

88 

89 return SharedChatModel.model_validate(shared_chat) 

90 

91 async def update(self, share_id: str, db: Optional[AsyncSession] = None) -> Optional[SharedChatModel]: 

92 """ 

93 Re-snapshot: update the shared chat with the current state of the original chat. 

94 """ 

95 async with get_async_db_context(db) as db: 

96 from open_webui.models.chats import Chat 

97 

98 shared_chat = await db.get(SharedChat, share_id) 

99 if not shared_chat: 

100 return None 

101 

102 chat = await db.get(Chat, shared_chat.chat_id) 

103 if not chat: 

104 return None 

105 

106 shared_chat.title = chat.title 

107 shared_chat.chat = chat.chat 

108 shared_chat.updated_at = int(time.time()) 

109 

110 await db.commit() 

111 await db.refresh(shared_chat) 

112 return SharedChatModel.model_validate(shared_chat) 

113 

114 async def get_by_id(self, share_id: str, db: Optional[AsyncSession] = None) -> Optional[SharedChatModel]: 

115 """Get a shared chat by its share token.""" 

116 async with get_async_db_context(db) as db: 

117 shared_chat = await db.get(SharedChat, share_id) 

118 if shared_chat: 

119 return SharedChatModel.model_validate(shared_chat) 

120 return None 

121 

122 async def get_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> Optional[SharedChatModel]: 

123 """Get the shared chat for a given original chat. Returns the most recent one.""" 

124 async with get_async_db_context(db) as db: 

125 result = await db.execute( 

126 select(SharedChat).filter_by(chat_id=chat_id).order_by(SharedChat.updated_at.desc()).limit(1) 

127 ) 

128 shared_chat = result.scalars().first() 

129 if shared_chat: 

130 return SharedChatModel.model_validate(shared_chat) 

131 return None 

132 

133 async def get_by_user_id( 

134 self, 

135 user_id: str, 

136 filter: Optional[dict] = None, 

137 skip: int = 0, 

138 limit: int = 50, 

139 db: Optional[AsyncSession] = None, 

140 ) -> list[SharedChatResponse]: 

141 """List all shared chats created by a user.""" 

142 async with get_async_db_context(db) as db: 

143 stmt = select(SharedChat).filter_by(user_id=user_id) 

144 

145 if filter: 

146 query_key = filter.get('query') 

147 if query_key: 

148 stmt = stmt.filter(SharedChat.title.ilike(f'%{query_key}%')) 

149 

150 order_by = filter.get('order_by') 

151 direction = filter.get('direction') 

152 

153 if order_by and direction: 

154 col = getattr(SharedChat, order_by, None) 

155 if not col: 155 ↛ 157line 155 didn't jump to line 157 because the condition on line 155 was always true

156 raise ValueError('Invalid order_by field') 

157 if direction.lower() == 'asc': 

158 stmt = stmt.order_by(col.asc()) 

159 elif direction.lower() == 'desc': 

160 stmt = stmt.order_by(col.desc()) 

161 else: 

162 raise ValueError('Invalid direction for ordering') 

163 else: 

164 stmt = stmt.order_by(SharedChat.updated_at.desc()) 

165 

166 if skip: 

167 stmt = stmt.offset(skip) 

168 if limit: 168 ↛ 171line 168 didn't jump to line 171 because the condition on line 168 was always true

169 stmt = stmt.limit(limit) 

170 

171 result = await db.execute(stmt) 

172 return [ 

173 SharedChatResponse( 

174 id=sc.chat_id, 

175 chat_id=sc.chat_id, 

176 title=sc.title, 

177 share_id=sc.id, 

178 updated_at=sc.updated_at, 

179 created_at=sc.created_at, 

180 ) 

181 for sc in result.scalars().all() 

182 ] 

183 

184 async def delete_by_id(self, share_id: str, db: Optional[AsyncSession] = None) -> bool: 

185 """Delete a shared chat by its share token.""" 

186 try: 

187 async with get_async_db_context(db) as db: 

188 await db.execute(delete(SharedChat).filter_by(id=share_id)) 

189 await db.commit() 

190 return True 

191 except Exception: 

192 return False 

193 

194 async def delete_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> bool: 

195 """Delete all shared chats for a given original chat.""" 

196 try: 

197 async with get_async_db_context(db) as db: 

198 await db.execute(delete(SharedChat).filter_by(chat_id=chat_id)) 

199 await db.commit() 

200 return True 

201 except Exception: 

202 return False 

203 

204 async def delete_all_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: 

205 """Delete all shared chats created by a user.""" 

206 try: 

207 async with get_async_db_context(db) as db: 

208 await db.execute(delete(SharedChat).filter_by(user_id=user_id)) 

209 await db.commit() 

210 return True 

211 except Exception: 

212 return False 

213 

214 

215SharedChats = SharedChatsTable()