add support for rds iam authentication

add support for rds iam authentication

add support for rds iam authentication

add support for rds iam authentication

fix arg name

rm async engine import, return type annotation

fix Path coercion of input ca path from env var

handle iam token auth prior to peewee migration

pass ssl args to peewee migration connection

use sslmode without cert verification if ca path can't be resolved

debug log

log level info -> warning

use rds iam config engine for alembic migration when using iam auth

rm debug log statements

fix formatting

add support for rds iam authentication

add support for rds iam authentication

fix arg name

rm async engine import, return type annotation

fix Path coercion of input ca path from env var

handle iam token auth prior to peewee migration

pass ssl args to peewee migration connection

use sslmode without cert verification if ca path can't be resolved

debug log

log level info -> warning

use rds iam config engine for alembic migration when using iam auth

rm debug log statements

fix formatting
This commit is contained in:
Brendan Shanahan 2026-04-08 17:03:33 -04:00
parent 1dcbfd47fb
commit ca61f5ae8c
4 changed files with 155 additions and 7 deletions

View file

@ -335,6 +335,10 @@ elif DATABASE_TYPE == 'sqlite+sqlcipher' and not os.environ.get('DATABASE_URL'):
if 'postgres://' in DATABASE_URL:
DATABASE_URL = DATABASE_URL.replace('postgres://', 'postgresql://')
DATABASE_ENABLE_IAM_TOKEN_AUTH = os.environ.get('DATABASE_ENABLE_IAM_TOKEN_AUTH', 'False').lower() == 'true'
DATABASE_CA_PATH = os.getenv('DATABASE_CA_PATH', None)
DATABASE_SCHEMA = os.environ.get('DATABASE_SCHEMA', None)
DATABASE_POOL_SIZE = os.environ.get('DATABASE_POOL_SIZE', None)

View file

@ -3,11 +3,19 @@ import json
import logging
from contextlib import contextmanager
from typing import Any, Optional
from pathlib import Path
from datetime import datetime, timedelta, timezone
from urllib.parse import parse_qs, quote_plus, urlsplit
import boto3
from open_webui.internal.wrappers import register_connection
from open_webui.env import (
OPEN_WEBUI_DIR,
DB_VARS,
DATABASE_CA_PATH,
DATABASE_ENABLE_IAM_TOKEN_AUTH,
DATABASE_URL,
DATABASE_USER,
DATABASE_SCHEMA,
DATABASE_POOL_MAX_OVERFLOW,
DATABASE_POOL_RECYCLE,
@ -18,7 +26,10 @@ from open_webui.env import (
ENABLE_DB_MIGRATIONS,
)
from peewee_migrate import Router
from pydantic import BaseModel, SecretStr
from sqlalchemy import Dialect, create_engine, MetaData, event, types
from sqlalchemy.engine import Engine
from sqlalchemy.engine.url import URL
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import scoped_session, sessionmaker, Session
from sqlalchemy.pool import QueuePool, NullPool
@ -28,6 +39,107 @@ from typing_extensions import Self
log = logging.getLogger(__name__)
class IAMToken(BaseModel):
token: SecretStr
expiration: datetime
class RDSIAMConfig:
def __init__(
self, db_type: str, db_name: str, db_host: str, db_port: int, db_user: str, ca_path: str | None = None
) -> None:
self.url = URL.create(
drivername=db_type,
username=db_user,
password=None,
host=db_host,
port=db_port,
database=db_name,
)
self.endpoint = db_host
self.port = db_port
self.db_user = db_user
self.connect_args = {"sslmode": "require"}
if ca_path is None:
log.info("No CA provided; using sslmode=require without certificate verification.")
elif isinstance(ca_path, str):
_ca_path = Path(ca_path).resolve()
if _ca_path.exists() and _ca_path.is_file():
self.connect_args["sslmode"] = "verify-full"
self.connect_args["sslrootcert"] = str(_ca_path)
else:
log.warning(f"CA file not found at {_ca_path}; using sslmode=require without certificate verification.")
else:
log.warning("Unable to verify CA file path; using sslmode=require without certificate verification.")
self.client = boto3.client("rds", region_name=os.getenv("AWS_DEFAULT_REGION", "us-east-1"))
self.token = self.get_token()
self.engine = self.create_engine()
def get_token(self) -> IAMToken:
"""Generate and return an RDS IAM authorization token.
RDS tokens are presigned URLs; expiration time is calculated from the token's
query string parameters.
"""
try:
token = self.client.generate_db_auth_token(self.endpoint, self.port, self.db_user)
expiration = self.set_token_expiration(token)
log.info(f"AWS RDS token for {self.endpoint} expires {expiration} UTC.")
return IAMToken(token=SecretStr(token), expiration=expiration)
except Exception as e:
log.error(f"AWS RDS token error: {e.__class__.__qualname__} {e}")
raise
def set_token_expiration(self, token: str) -> datetime:
"""
Calculate authorization token expiration time from a presigned RDS URL.
RDS tokens are presigned URLs. The query string contains the time at which the
token was generated, and the number of seconds until it expires. This function
converts these values and returns the token expiration as a UTC-aware datetime.
"""
split_token = urlsplit(token)
parsed_qs: dict = parse_qs(split_token.query)
generation_time = datetime.strptime(parsed_qs["X-Amz-Date"][0], "%Y%m%dT%H%M%SZ")
duration = timedelta(seconds=int(parsed_qs["X-Amz-Expires"][0]))
expiration = generation_time + duration
return expiration.replace(tzinfo=timezone.utc)
def create_engine(self) -> Engine:
if isinstance(DATABASE_POOL_SIZE, int):
if DATABASE_POOL_SIZE > 0:
engine = create_engine(
self.url,
pool_size=DATABASE_POOL_SIZE,
max_overflow=DATABASE_POOL_MAX_OVERFLOW,
pool_timeout=DATABASE_POOL_TIMEOUT,
pool_recycle=DATABASE_POOL_RECYCLE,
pool_pre_ping=True,
poolclass=QueuePool,
connect_args=self.connect_args,
)
else:
engine = create_engine(
self.url,
pool_pre_ping=True,
poolclass=NullPool,
connect_args=self.connect_args,
)
else:
engine = create_engine(self.url, pool_pre_ping=True, connect_args=self.connect_args)
return engine
def check_token(self, dialect, connection_record, connection_args, connection_kwargs) -> None:
now = datetime.now(tz=timezone.utc)
if not self.token or now >= self.token.expiration:
self.token = self.get_token()
connection_kwargs["password"] = self.token.token.get_secret_value()
class JSONField(types.TypeDecorator):
impl = types.Text
cache_ok = True
@ -52,11 +164,11 @@ class JSONField(types.TypeDecorator):
# Workaround to handle the peewee migration
# This is required to ensure the peewee migration is handled before the alembic migration
def handle_peewee_migration(DATABASE_URL):
# db = None
def handle_peewee_migration(db_url: str = DATABASE_URL, connect_args: dict | None = None):
db = None
try:
# Replace the postgresql:// with postgres:// to handle the peewee migration
db = register_connection(DATABASE_URL.replace('postgresql://', 'postgres://'))
db = register_connection(db_url.replace('postgresql://', 'postgres://'), connect_args=connect_args)
migrate_dir = OPEN_WEBUI_DIR / 'internal' / 'migrations'
router = Router(db, logger=log, migrate_dir=migrate_dir)
router.run()
@ -72,11 +184,30 @@ def handle_peewee_migration(DATABASE_URL):
db.close()
# Assert if db connection has been closed
assert db.is_closed(), 'Database connection is still open.'
if db is not None:
assert db.is_closed(), 'Database connection is still open.'
if DATABASE_ENABLE_IAM_TOKEN_AUTH:
rds_iam_config = RDSIAMConfig(
db_type=DB_VARS["db_type"],
db_name=DB_VARS["db_name"],
db_host=DB_VARS["db_host"],
db_port=DB_VARS["db_port"],
db_user=DATABASE_USER,
ca_path=DATABASE_CA_PATH,
)
if ENABLE_DB_MIGRATIONS:
handle_peewee_migration(DATABASE_URL)
if DATABASE_ENABLE_IAM_TOKEN_AUTH:
_token_val = rds_iam_config.token.token.get_secret_value()
_migration_url = (
f'{DB_VARS["db_type"]}://{DATABASE_USER}:{quote_plus(_token_val)}'
f'@{DB_VARS["db_host"]}:{DB_VARS["db_port"]}/{DB_VARS["db_name"]}'
)
handle_peewee_migration(_migration_url, connect_args=rds_iam_config.connect_args)
else:
handle_peewee_migration(DATABASE_URL)
SQLALCHEMY_DATABASE_URL = DATABASE_URL
@ -137,6 +268,9 @@ elif 'sqlite' in SQLALCHEMY_DATABASE_URL:
cursor.close()
event.listen(engine, 'connect', on_connect)
elif DATABASE_ENABLE_IAM_TOKEN_AUTH:
engine = rds_iam_config.engine
event.listen(engine, 'do_connect', rds_iam_config.check_token)
else:
if isinstance(DATABASE_POOL_SIZE, int):
if DATABASE_POOL_SIZE > 0:

View file

@ -41,7 +41,7 @@ class ReconnectingPostgresqlDatabase(CustomReconnectMixin, PostgresqlDatabase):
pass
def register_connection(db_url):
def register_connection(db_url: str, connect_args: dict | None = None):
# Check if using SQLCipher protocol
if db_url.startswith('sqlite+sqlcipher://'):
database_password = os.environ.get('DATABASE_PASSWORD')
@ -71,7 +71,14 @@ def register_connection(db_url):
# Get the connection details
connection = parse(db_url, unquote_user=True, unquote_password=True)
if connect_args:
if "sslmode" in connect_args:
connection["sslmode"] = connect_args["sslmode"]
if "sslrootcert" in connect_args:
connection["sslrootcert"] = connect_args["sslrootcert"]
# Use our custom database class that supports reconnection
log.debug(f"Peewee connection params (password redacted): { {k: ('***' if k == 'password' else v) for k, v in connection.items()} }")
db = ReconnectingPostgresqlDatabase(**connection)
db.connect(reuse_if_open=True)
elif isinstance(db, SqliteDatabase):

View file

@ -3,7 +3,7 @@ from logging.config import fileConfig
from alembic import context
from open_webui.models.auths import Auth
from open_webui.env import DATABASE_URL, DATABASE_PASSWORD, LOG_FORMAT
from open_webui.env import DATABASE_URL, DATABASE_ENABLE_IAM_TOKEN_AUTH, DATABASE_PASSWORD, LOG_FORMAT
from sqlalchemy import engine_from_config, pool, create_engine
# this is the Alembic Config object, which provides
@ -93,6 +93,9 @@ def run_migrations_online() -> None:
creator=create_sqlcipher_connection,
echo=False,
)
elif DATABASE_ENABLE_IAM_TOKEN_AUTH:
from open_webui.internal.db import rds_iam_config
connectable = rds_iam_config.engine
else:
# Standard database connection (existing logic)
connectable = engine_from_config(