|
69 | 69 | ENABLE_OAUTH_ID_TOKEN_COOKIE, |
70 | 70 | ENABLE_OAUTH_EMAIL_FALLBACK, |
71 | 71 | OAUTH_CLIENT_INFO_ENCRYPTION_KEY, |
| 72 | + OAUTH_MAX_SESSIONS_PER_USER, |
72 | 73 | ) |
73 | 74 | from open_webui.utils.misc import parse_duration |
74 | 75 | from open_webui.utils.auth import get_password_hash, create_token |
@@ -1679,11 +1680,23 @@ async def handle_callback(self, request, provider, response, db=None): |
1679 | 1680 | if "expires_in" in token and "expires_at" not in token: |
1680 | 1681 | token["expires_at"] = datetime.now().timestamp() + token["expires_in"] |
1681 | 1682 |
|
1682 | | - # Clean up any existing sessions for this user/provider first |
| 1683 | + # Enforce max concurrent sessions per user/provider to prevent |
| 1684 | + # unbounded growth while allowing multi-device usage |
1683 | 1685 | sessions = OAuthSessions.get_sessions_by_user_id(user.id, db=db) |
1684 | | - for session in sessions: |
1685 | | - if session.provider == provider: |
1686 | | - OAuthSessions.delete_session_by_id(session.id, db=db) |
| 1686 | + provider_sessions = sorted( |
| 1687 | + [ |
| 1688 | + for session in sessions |
| 1689 | + if session.provider == provider |
| 1690 | + ], |
| 1691 | + key=lambda session: session.created_at, |
| 1692 | + reverse=True, |
| 1693 | + ) |
| 1694 | + # Keep the newest sessions up to the limit, prune the rest |
| 1695 | + if len(provider_sessions) >= OAUTH_MAX_SESSIONS_PER_USER: |
| 1696 | + for old_session in provider_sessions[ |
| 1697 | + OAUTH_MAX_SESSIONS_PER_USER - 1 : |
| 1698 | + ]: |
| 1699 | + OAuthSessions.delete_session_by_id(old_session.id, db=db) |
1687 | 1700 |
|
1688 | 1701 | session = OAuthSessions.create_session( |
1689 | 1702 | user_id=user.id, |
|
0 commit comments