mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: 扩展多个字符串列为 TEXT 类型并添加 tls-client 可选依赖
数据库迁移: - 扩展 provider_api_keys.api_key 为 TEXT(OAuth tokens 可能很长) - 扩展 ldap_configs 的 bind_dn、base_dn、user_search_filter 为 TEXT - 扩展 oauth_providers.client_id 为 TEXT - 添加 SQLite 兼容支持(batch 模式) - 添加表/列存在性检查 依赖: - 添加 tls-client 作为可选依赖 [tls]
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
"""Add provider_type and expand api_key column to TEXT
|
||||
"""Add provider_type and expand string columns to TEXT
|
||||
|
||||
- Add providers.provider_type (String(20), server_default="custom")
|
||||
- Change provider_api_keys.api_key from VARCHAR(500) to TEXT (OAuth tokens can be long)
|
||||
- Expand multiple VARCHAR columns to TEXT for long values (OAuth tokens, LDAP DN, URLs, etc.)
|
||||
|
||||
Revision ID: b5c6d7e8f9a0
|
||||
Revises: c4e8f9a1b2c3
|
||||
@@ -22,6 +22,16 @@ branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
# 需要扩展为 TEXT 的列(表名, 列名, 原始类型长度)
|
||||
COLUMNS_TO_EXPAND = [
|
||||
("provider_api_keys", "api_key", 500), # OAuth tokens can be very long
|
||||
("ldap_configs", "bind_dn", 255), # LDAP DN can be deeply nested
|
||||
("ldap_configs", "base_dn", 255), # LDAP DN can be deeply nested
|
||||
("ldap_configs", "user_search_filter", 500), # Complex LDAP filters
|
||||
("oauth_providers", "client_id", 255), # Some OAuth providers use JWT client_id
|
||||
]
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
"""检查列是否已存在"""
|
||||
bind = op.get_bind()
|
||||
@@ -30,6 +40,72 @@ def column_exists(table_name: str, column_name: str) -> bool:
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
"""检查表是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def is_sqlite() -> bool:
|
||||
"""检查是否为 SQLite 数据库"""
|
||||
bind = op.get_bind()
|
||||
return bind.dialect.name == "sqlite"
|
||||
|
||||
|
||||
def expand_column_to_text(table_name: str, column_name: str, original_length: int) -> None:
|
||||
"""将 VARCHAR 列扩展为 TEXT(兼容 SQLite)"""
|
||||
if not table_exists(table_name):
|
||||
return
|
||||
if not column_exists(table_name, column_name):
|
||||
return
|
||||
|
||||
if is_sqlite():
|
||||
# SQLite 不支持直接 ALTER COLUMN,需要用 batch 模式
|
||||
with op.batch_alter_table(table_name) as batch_op:
|
||||
batch_op.alter_column(
|
||||
column_name,
|
||||
type_=sa.Text(),
|
||||
existing_type=sa.String(original_length),
|
||||
)
|
||||
else:
|
||||
op.alter_column(
|
||||
table_name,
|
||||
column_name,
|
||||
type_=sa.Text(),
|
||||
existing_type=sa.String(original_length),
|
||||
)
|
||||
|
||||
|
||||
def shrink_column_to_varchar(
|
||||
table_name: str, column_name: str, target_length: int, nullable: bool = False
|
||||
) -> None:
|
||||
"""将 TEXT 列缩小为 VARCHAR(兼容 SQLite)
|
||||
WARNING: 如果数据超过 target_length 会失败
|
||||
"""
|
||||
if not table_exists(table_name):
|
||||
return
|
||||
if not column_exists(table_name, column_name):
|
||||
return
|
||||
|
||||
if is_sqlite():
|
||||
with op.batch_alter_table(table_name) as batch_op:
|
||||
batch_op.alter_column(
|
||||
column_name,
|
||||
type_=sa.String(target_length),
|
||||
existing_type=sa.Text(),
|
||||
existing_nullable=nullable,
|
||||
)
|
||||
else:
|
||||
op.alter_column(
|
||||
table_name,
|
||||
column_name,
|
||||
type_=sa.String(target_length),
|
||||
existing_type=sa.Text(),
|
||||
existing_nullable=nullable,
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Add providers.provider_type
|
||||
if not column_exists("providers", "provider_type"):
|
||||
@@ -38,26 +114,16 @@ def upgrade() -> None:
|
||||
sa.Column("provider_type", sa.String(20), nullable=False, server_default="custom"),
|
||||
)
|
||||
|
||||
# Expand provider_api_keys.api_key from VARCHAR(500) to TEXT
|
||||
op.alter_column(
|
||||
"provider_api_keys",
|
||||
"api_key",
|
||||
type_=sa.Text(),
|
||||
existing_type=sa.String(500),
|
||||
existing_nullable=False,
|
||||
)
|
||||
# Expand VARCHAR columns to TEXT
|
||||
for table_name, column_name, original_length in COLUMNS_TO_EXPAND:
|
||||
expand_column_to_text(table_name, column_name, original_length)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Revert provider_api_keys.api_key from TEXT to VARCHAR(500)
|
||||
# WARNING: Downgrade may fail if any api_key values exceed 500 characters
|
||||
op.alter_column(
|
||||
"provider_api_keys",
|
||||
"api_key",
|
||||
type_=sa.String(500),
|
||||
existing_type=sa.Text(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
# Shrink TEXT columns back to VARCHAR
|
||||
# WARNING: Downgrade may fail if any values exceed original length
|
||||
for table_name, column_name, original_length in reversed(COLUMNS_TO_EXPAND):
|
||||
shrink_column_to_varchar(table_name, column_name, original_length)
|
||||
|
||||
# Drop providers.provider_type
|
||||
if column_exists("providers", "provider_type"):
|
||||
|
||||
@@ -482,12 +482,12 @@ class LDAPConfig(Base):
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
server_url = Column(String(255), nullable=False) # ldap://host:389 或 ldaps://host:636
|
||||
bind_dn = Column(String(255), nullable=False) # 绑定账号 DN
|
||||
bind_dn = Column(Text, nullable=False) # 绑定账号 DN(可能很长)
|
||||
bind_password_encrypted = Column(Text, nullable=True) # 加密的绑定密码(允许 NULL 表示已清除)
|
||||
base_dn = Column(String(255), nullable=False) # 用户搜索基础 DN
|
||||
base_dn = Column(Text, nullable=False) # 用户搜索基础 DN(可能很长)
|
||||
user_search_filter = Column(
|
||||
String(500), default="(uid={username})", nullable=False
|
||||
) # 用户搜索过滤器
|
||||
Text, default="(uid={username})", nullable=False
|
||||
) # 用户搜索过滤器(可能很复杂)
|
||||
username_attr = Column(
|
||||
String(50), default="uid", nullable=False
|
||||
) # 用户名属性 (uid/sAMAccountName)
|
||||
@@ -548,7 +548,7 @@ class OAuthProvider(Base):
|
||||
provider_type = Column(String(50), primary_key=True)
|
||||
display_name = Column(String(100), nullable=False)
|
||||
|
||||
client_id = Column(String(255), nullable=False)
|
||||
client_id = Column(Text, nullable=False) # 某些 OAuth 提供商可能使用很长的 client_id
|
||||
client_secret_encrypted = Column(Text, nullable=True) # 允许 NULL 表示尚未配置/已清除
|
||||
|
||||
# 可选覆盖端点(需在业务层做白名单校验)
|
||||
|
||||
15
uv.lock
generated
15
uv.lock
generated
@@ -40,6 +40,9 @@ dev = [
|
||||
{ name = "pytest" },
|
||||
{ name = "pytest-asyncio" },
|
||||
]
|
||||
tls = [
|
||||
{ name = "tls-client" },
|
||||
]
|
||||
|
||||
[package.dev-dependencies]
|
||||
dev = [
|
||||
@@ -86,9 +89,10 @@ requires-dist = [
|
||||
{ name = "regex", specifier = ">=2026.1.15" },
|
||||
{ name = "sqlalchemy", specifier = ">=2.0.46" },
|
||||
{ name = "tiktoken", specifier = ">=0.12.0" },
|
||||
{ name = "tls-client", marker = "extra == 'tls'", specifier = ">=1.0.1" },
|
||||
{ name = "uvicorn", specifier = ">=0.40.0" },
|
||||
]
|
||||
provides-extras = ["dev"]
|
||||
provides-extras = ["dev", "tls"]
|
||||
|
||||
[package.metadata.requires-dev]
|
||||
dev = [
|
||||
@@ -2509,6 +2513,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/af/df/c7891ef9d2712ad774777271d39fdef63941ffba0a9d59b7ad1fd2765e57/tiktoken-0.12.0-cp314-cp314t-win_amd64.whl", hash = "sha256:f61c0aea5565ac82e2ec50a05e02a6c44734e91b51c10510b084ea1b8e633a71", size = 920667 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tls-client"
|
||||
version = "1.0.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/3c/a6/6ec27c66a836a11a085e841f825d2f5bd289092f3bcd2f645558f587c89f/tls_client-1.0.1.tar.gz", hash = "sha256:dad797f3412bb713606e0765d489f547ffb580c5ffdb74aed47a183ce8505ff5", size = 16414 }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/75/cd/5c735818692927e07980357445569adb6ee204c3332d19c516bae01c6cfa/tls_client-1.0.1-py3-none-any.whl", hash = "sha256:2f8915c0642c2226c9e33120072a2af082812f6310d32f4ea4da322db7d3bb1c", size = 41287556 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tqdm"
|
||||
version = "4.67.1"
|
||||
|
||||
Reference in New Issue
Block a user