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:
fawney19
2026-02-04 16:45:05 +08:00
parent c996078f30
commit 24c9105628
3 changed files with 104 additions and 25 deletions

View File

@@ -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"):

View File

@@ -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
View File

@@ -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"