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
« 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
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
11log = logging.getLogger(__name__)
13####################
14# SharedChat DB Schema
15####################
18class SharedChat(Base):
19 __tablename__ = 'shared_chat'
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
25 title = Column(Text)
26 chat = Column(JSON) # Snapshot of chat JSON at share time
28 created_at = Column(BigInteger)
29 updated_at = Column(BigInteger)
32class SharedChatModel(BaseModel):
33 model_config = ConfigDict(from_attributes=True)
35 id: str
36 chat_id: str
37 user_id: str
39 title: str
40 chat: dict
42 created_at: int
43 updated_at: int
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
55####################
56# Table Operations
57####################
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
69 chat = await db.get(Chat, chat_id)
70 if not chat:
71 return None
73 share_id = str(uuid.uuid4())
74 now = int(time.time())
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)
89 return SharedChatModel.model_validate(shared_chat)
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
98 shared_chat = await db.get(SharedChat, share_id)
99 if not shared_chat:
100 return None
102 chat = await db.get(Chat, shared_chat.chat_id)
103 if not chat:
104 return None
106 shared_chat.title = chat.title
107 shared_chat.chat = chat.chat
108 shared_chat.updated_at = int(time.time())
110 await db.commit()
111 await db.refresh(shared_chat)
112 return SharedChatModel.model_validate(shared_chat)
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
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
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)
145 if filter:
146 query_key = filter.get('query')
147 if query_key:
148 stmt = stmt.filter(SharedChat.title.ilike(f'%{query_key}%'))
150 order_by = filter.get('order_by')
151 direction = filter.get('direction')
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())
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)
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 ]
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
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
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
215SharedChats = SharedChatsTable()