Coverage for open_webui/models/oauth_sessions.py: 35%
183 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 base64
2import hashlib
3import logging
4import time
5import uuid
6from typing import List, Optional
8from cryptography.fernet import Fernet
9from open_webui.env import OAUTH_SESSION_TOKEN_ENCRYPTION_KEY
10from open_webui.internal.db import Base, get_async_db_context
11from open_webui.utils.json_codec import JSONCodec
12from pydantic import BaseModel, ConfigDict
13from sqlalchemy import BigInteger, Column, Index, String, Text, delete, select, update
14from sqlalchemy.ext.asyncio import AsyncSession
16log = logging.getLogger(__name__)
18####################
19# DB MODEL
20####################
23class OAuthSession(Base):
24 __tablename__ = 'oauth_session'
26 id = Column(Text, primary_key=True, unique=True)
27 user_id = Column(Text, nullable=False)
28 provider = Column(Text, nullable=False)
29 token = Column(Text, nullable=False) # JSON with access_token, id_token, refresh_token
30 expires_at = Column(BigInteger, nullable=False)
31 created_at = Column(BigInteger, nullable=False)
32 updated_at = Column(BigInteger, nullable=False)
34 # Add indexes for better performance
35 __table_args__ = (
36 Index('idx_oauth_session_user_id', 'user_id'),
37 Index('idx_oauth_session_expires_at', 'expires_at'),
38 Index('idx_oauth_session_user_provider', 'user_id', 'provider'),
39 )
42class OAuthSessionModel(BaseModel):
43 id: str
44 user_id: str
45 provider: str
46 token: dict
47 expires_at: int # timestamp in epoch
48 created_at: int # timestamp in epoch
49 updated_at: int # timestamp in epoch
51 model_config = ConfigDict(from_attributes=True)
54####################
55# Forms
56####################
59class OAuthSessionResponse(BaseModel):
60 id: str
61 user_id: str
62 provider: str
63 expires_at: int
66class OAuthSessionTable:
67 def __init__(self):
68 self.encryption_key = OAUTH_SESSION_TOKEN_ENCRYPTION_KEY
69 if not self.encryption_key: 69 ↛ 70line 69 didn't jump to line 70 because the condition on line 69 was never true
70 raise Exception('OAUTH_SESSION_TOKEN_ENCRYPTION_KEY is not set')
72 # check if encryption key is in the right format for Fernet (32 url-safe base64-encoded bytes)
73 if len(self.encryption_key) != 44: 73 ↛ 77line 73 didn't jump to line 77 because the condition on line 73 was always true
74 key_bytes = hashlib.sha256(self.encryption_key.encode()).digest()
75 self.encryption_key = base64.urlsafe_b64encode(key_bytes)
76 else:
77 self.encryption_key = self.encryption_key.encode()
79 try:
80 self.fernet = Fernet(self.encryption_key)
81 except Exception as e:
82 log.error(f'Error initializing Fernet with provided key: {e}')
83 raise
85 def _encrypt_token(self, token) -> str:
86 """Encrypt OAuth tokens for storage"""
87 try:
88 token_json = JSONCodec.dumps(token)
89 encrypted = self.fernet.encrypt(token_json.encode()).decode()
90 return encrypted
91 except Exception as e:
92 log.error(f'Error encrypting tokens: {e}')
93 raise
95 def _decrypt_token(self, token: str):
96 """Decrypt OAuth tokens from storage"""
97 try:
98 decrypted = self.fernet.decrypt(token.encode()).decode()
99 return JSONCodec.loads(decrypted)
100 except Exception as e:
101 log.error(f'Error decrypting tokens: {type(e).__name__}: {e}')
102 raise
104 async def create_session(
105 self,
106 user_id: str,
107 provider: str,
108 token: dict,
109 db: Optional[AsyncSession] = None,
110 ) -> Optional[OAuthSessionModel]:
111 """Create a new OAuth session"""
112 try:
113 async with get_async_db_context(db) as db:
114 current_time = int(time.time())
115 id = str(uuid.uuid4())
117 result = OAuthSession(
118 **{
119 'id': id,
120 'user_id': user_id,
121 'provider': provider,
122 'token': self._encrypt_token(token),
123 'expires_at': token.get('expires_at') or int(time.time() + 3600),
124 'created_at': current_time,
125 'updated_at': current_time,
126 }
127 )
129 db.add(result)
130 await db.commit()
132 if result:
133 # Make a copy of the model data before closing session
134 model = OAuthSessionModel(
135 id=result.id,
136 user_id=result.user_id,
137 provider=result.provider,
138 token=token, # Return decrypted token
139 expires_at=result.expires_at,
140 created_at=result.created_at,
141 updated_at=result.updated_at,
142 )
143 return model
144 else:
145 return None
146 except Exception as e:
147 log.error(f'Error creating OAuth session: {e}')
148 return None
150 async def get_session_by_id(
151 self, session_id: str, db: Optional[AsyncSession] = None
152 ) -> Optional[OAuthSessionModel]:
153 """Get OAuth session by ID"""
154 try:
155 async with get_async_db_context(db) as db:
156 result = await db.execute(select(OAuthSession).filter_by(id=session_id))
157 session = result.scalars().first()
158 if session:
159 return OAuthSessionModel(
160 id=session.id,
161 user_id=session.user_id,
162 provider=session.provider,
163 token=self._decrypt_token(session.token),
164 expires_at=session.expires_at,
165 created_at=session.created_at,
166 updated_at=session.updated_at,
167 )
169 return None
170 except Exception as e:
171 log.error(f'Error getting OAuth session by ID: {e}')
172 return None
174 async def get_session_by_id_and_user_id(
175 self, session_id: str, user_id: str, db: Optional[AsyncSession] = None
176 ) -> Optional[OAuthSessionModel]:
177 """Get OAuth session by ID and user ID"""
178 try:
179 async with get_async_db_context(db) as db:
180 result = await db.execute(select(OAuthSession).filter_by(id=session_id, user_id=user_id))
181 session = result.scalars().first()
182 if session:
183 return OAuthSessionModel(
184 id=session.id,
185 user_id=session.user_id,
186 provider=session.provider,
187 token=self._decrypt_token(session.token),
188 expires_at=session.expires_at,
189 created_at=session.created_at,
190 updated_at=session.updated_at,
191 )
193 return None
194 except Exception as e:
195 log.error(f'Error getting OAuth session by ID: {e}')
196 return None
198 async def get_session_by_provider_and_user_id(
199 self, provider: str, user_id: str, db: Optional[AsyncSession] = None
200 ) -> Optional[OAuthSessionModel]:
201 """Get OAuth session by provider and user ID"""
202 try:
203 async with get_async_db_context(db) as db:
204 result = await db.execute(
205 select(OAuthSession)
206 .filter_by(provider=provider, user_id=user_id)
207 .order_by(OAuthSession.created_at.desc())
208 )
209 session = result.scalars().first()
210 if session:
211 return OAuthSessionModel(
212 id=session.id,
213 user_id=session.user_id,
214 provider=session.provider,
215 token=self._decrypt_token(session.token),
216 expires_at=session.expires_at,
217 created_at=session.created_at,
218 updated_at=session.updated_at,
219 )
221 return None
222 except Exception as e:
223 log.error(f'Error getting OAuth session by provider and user ID: {e}')
224 return None
226 async def get_sessions_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> List[OAuthSessionModel]:
227 """Get all OAuth sessions for a user"""
228 try:
229 async with get_async_db_context(db) as db:
230 result = await db.execute(select(OAuthSession).filter_by(user_id=user_id))
231 sessions = result.scalars().all()
233 results = []
234 for session in sessions:
235 try:
236 results.append(
237 OAuthSessionModel(
238 id=session.id,
239 user_id=session.user_id,
240 provider=session.provider,
241 token=self._decrypt_token(session.token),
242 expires_at=session.expires_at,
243 created_at=session.created_at,
244 updated_at=session.updated_at,
245 )
246 )
247 except Exception as e:
248 log.warning(
249 f'Skipping OAuth session {session.id} due to decryption failure, deleting corrupted session: {type(e).__name__}: {e}'
250 )
251 await db.execute(delete(OAuthSession).filter_by(id=session.id))
252 await db.commit()
254 return results
256 except Exception as e:
257 log.error(f'Error getting OAuth sessions by user ID: {e}')
258 return []
260 async def update_session_by_id(
261 self, session_id: str, token: dict, db: Optional[AsyncSession] = None
262 ) -> Optional[OAuthSessionModel]:
263 """Update OAuth session tokens"""
264 try:
265 async with get_async_db_context(db) as db:
266 current_time = int(time.time())
268 await db.execute(
269 update(OAuthSession)
270 .filter_by(id=session_id)
271 .values(
272 token=self._encrypt_token(token),
273 expires_at=token.get('expires_at') or int(time.time() + 3600),
274 updated_at=current_time,
275 )
276 )
277 await db.commit()
278 result = await db.execute(select(OAuthSession).filter_by(id=session_id))
279 session = result.scalars().first()
281 if session:
282 return OAuthSessionModel(
283 id=session.id,
284 user_id=session.user_id,
285 provider=session.provider,
286 token=self._decrypt_token(session.token),
287 expires_at=session.expires_at,
288 created_at=session.created_at,
289 updated_at=session.updated_at,
290 )
292 return None
293 except Exception as e:
294 log.error(f'Error updating OAuth session tokens: {e}')
295 return None
297 async def delete_session_by_id(self, session_id: str, db: Optional[AsyncSession] = None) -> bool:
298 """Delete an OAuth session"""
299 try:
300 async with get_async_db_context(db) as db:
301 result = await db.execute(delete(OAuthSession).filter_by(id=session_id))
302 await db.commit()
303 return result.rowcount > 0
304 except Exception as e:
305 log.error(f'Error deleting OAuth session: {e}')
306 return False
308 async def delete_sessions_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
309 """Delete all OAuth sessions for a user"""
310 try:
311 async with get_async_db_context(db) as db:
312 await db.execute(delete(OAuthSession).filter_by(user_id=user_id))
313 await db.commit()
314 return True
315 except Exception as e:
316 log.error(f'Error deleting OAuth sessions by user ID: {e}')
317 return False
319 async def delete_sessions_by_user_id_and_provider(
320 self, user_id: str, provider: str, db: Optional[AsyncSession] = None
321 ) -> bool:
322 """Delete all OAuth sessions for a specific user and provider"""
323 try:
324 async with get_async_db_context(db) as db:
325 result = await db.execute(delete(OAuthSession).filter_by(user_id=user_id, provider=provider))
326 await db.commit()
327 return result.rowcount > 0
328 except Exception as e:
329 log.error(f'Error deleting OAuth sessions for user {user_id} and provider {provider}: {e}')
330 return False
332 async def delete_sessions_by_provider(self, provider: str, db: Optional[AsyncSession] = None) -> bool:
333 """Delete all OAuth sessions for a provider"""
334 try:
335 async with get_async_db_context(db) as db:
336 await db.execute(delete(OAuthSession).filter_by(provider=provider))
337 await db.commit()
338 return True
339 except Exception as e:
340 log.error(f'Error deleting OAuth sessions by provider {provider}: {e}')
341 return False
344OAuthSessions = OAuthSessionTable()