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

1import base64 

2import hashlib 

3import logging 

4import time 

5import uuid 

6from typing import List, Optional 

7 

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 

15 

16log = logging.getLogger(__name__) 

17 

18#################### 

19# DB MODEL 

20#################### 

21 

22 

23class OAuthSession(Base): 

24 __tablename__ = 'oauth_session' 

25 

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) 

33 

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 ) 

40 

41 

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 

50 

51 model_config = ConfigDict(from_attributes=True) 

52 

53 

54#################### 

55# Forms 

56#################### 

57 

58 

59class OAuthSessionResponse(BaseModel): 

60 id: str 

61 user_id: str 

62 provider: str 

63 expires_at: int 

64 

65 

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

71 

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

78 

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 

84 

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 

94 

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 

103 

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

116 

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 ) 

128 

129 db.add(result) 

130 await db.commit() 

131 

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 

149 

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 ) 

168 

169 return None 

170 except Exception as e: 

171 log.error(f'Error getting OAuth session by ID: {e}') 

172 return None 

173 

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 ) 

192 

193 return None 

194 except Exception as e: 

195 log.error(f'Error getting OAuth session by ID: {e}') 

196 return None 

197 

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 ) 

220 

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 

225 

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

232 

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

253 

254 return results 

255 

256 except Exception as e: 

257 log.error(f'Error getting OAuth sessions by user ID: {e}') 

258 return [] 

259 

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

267 

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

280 

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 ) 

291 

292 return None 

293 except Exception as e: 

294 log.error(f'Error updating OAuth session tokens: {e}') 

295 return None 

296 

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 

307 

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 

318 

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 

331 

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 

342 

343 

344OAuthSessions = OAuthSessionTable()