refactor: 优化代码格式与修复Snowflake schema配置问题

1. 合并多行异常抛出为单行简化代码
2. 修正Snowflake驱动配置中schema字段名冲突,使用schema_name作为内部字段并通过alias保留schema对外键名
3. 调整异步任务调用的换行格式提升可读性
4. 为Pydantic模型添加populate_by_name配置支持按名称填充字段
This commit is contained in:
Kris 2026-07-07 16:20:58 +08:00
parent 53e067aa4b
commit 56cffed4e0
5 changed files with 19 additions and 19 deletions

View File

@ -84,8 +84,8 @@ class SnowflakeDriver(BaseDatabaseDriver):
kwargs["warehouse"] = self._db_config.warehouse
if self._db_config.database:
kwargs["database"] = self._db_config.database
if getattr(self._db_config, "schema", ""):
kwargs["schema"] = self._db_config.schema
if getattr(self._db_config, "schema_name", ""):
kwargs["schema"] = self._db_config.schema_name
if getattr(self._db_config, "role", ""):
kwargs["role"] = self._db_config.role
return kwargs

View File

@ -218,8 +218,8 @@ class DatabaseExecutor:
overrides["warehouse"] = conn_config.warehouse
if "role" in raw and conn_config.role:
overrides["role"] = conn_config.role
if "schema" in raw and conn_config.schema:
overrides["schema"] = conn_config.schema
if "schema" in raw and conn_config.schema_name:
overrides["schema_name"] = conn_config.schema_name
if not overrides:
return db_config
return db_config.model_copy(update=overrides)
@ -254,7 +254,7 @@ class DatabaseExecutor:
if db_config.db_type == "snowflake":
return (
f"{slug}:{db_config.db_type}:{db_config.account}:"
f"{db_config.warehouse}:{db_config.database}:{db_config.schema}:{cred_hash}"
f"{db_config.warehouse}:{db_config.database}:{db_config.schema_name}:{cred_hash}"
)
return f"{slug}:{db_config.db_type}:{db_config.host}:{db_config.port}:{db_config.database}:{cred_hash}"

View File

@ -10,6 +10,8 @@ from pydantic import BaseModel, ConfigDict, Field, field_validator, model_valida
class DatabaseConnectionConfig(BaseModel):
"""数据库系统/环境级连接配置。"""
model_config = ConfigDict(populate_by_name=True)
db_type: Literal["mysql", "postgresql", "mongodb", "bigquery", "databricks", "snowflake"] = Field(
default="postgresql"
)
@ -26,7 +28,8 @@ class DatabaseConnectionConfig(BaseModel):
account: str = Field(default="", description="Snowflake 账户标识符(如 xy12345.us-east-1")
warehouse: str = Field(default="", description="Snowflake 仓库名(如 COMPUTE_WH")
role: str = Field(default="", description="Snowflake 角色名(如 AGENT_ROLE")
schema: str = Field(default="", description="Snowflake 模式名(如 PUBLIC")
# 字段名避开 BaseModel.schema(),通过 alias 保留外部 "schema" 键名
schema_name: str = Field(default="", alias="schema", description="Snowflake 模式名(如 PUBLIC")
class DatabaseParameter(BaseModel):
@ -52,7 +55,7 @@ class DatabaseParameter(BaseModel):
class DatabaseAdapterConfig(BaseModel):
"""数据库适配器专属配置。"""
model_config = ConfigDict(extra="forbid")
model_config = ConfigDict(extra="forbid", populate_by_name=True)
# 连接配置
db_type: Literal["mysql", "postgresql", "mongodb", "bigquery", "databricks", "snowflake"] = Field(
@ -83,7 +86,8 @@ class DatabaseAdapterConfig(BaseModel):
)
warehouse: str = Field(default="", description="Snowflake 仓库名(如 COMPUTE_WH")
role: str = Field(default="", description="Snowflake 角色名(如 AGENT_ROLE")
schema: str = Field(default="", description="Snowflake 模式名(如 PUBLIC")
# 字段名避开 BaseModel.schema(),通过 alias 保留外部 "schema" 键名
schema_name: str = Field(default="", alias="schema", description="Snowflake 模式名(如 PUBLIC")
# SSL/TLS
ssl_mode: Literal["disable", "prefer", "require", "verify-ca", "verify-full"] = Field(default="disable")

View File

@ -783,9 +783,7 @@ class ImportService(ImportServicePort):
return result
if isinstance(raw, list):
return [d if isinstance(d, dict) else d.model_dump() for d in raw]
raise DomainValidationError(
f"适配器 generate_from_asset 返回了不支持的类型 {type(raw).__name__}"
)
raise DomainValidationError(f"适配器 generate_from_asset 返回了不支持的类型 {type(raw).__name__}")
def _drafts_from_dicts(self, drafts: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""将前端传回的 draft dict 转换为规范化的工具草稿 dict。"""

View File

@ -379,14 +379,12 @@ class QuotaService(QuotaServicePort):
system_id = input_dto.system_id
now = datetime.now(UTC)
total, by_window, near_warning_count, near_critical_count, expiring_count = (
await asyncio.gather(
self._repos.quota_usage.count(system_id=system_id),
self._repos.quota_usage.count_by_quota_window(system_id=system_id),
self._repos.quota_usage.count(system_id=system_id, near_warning=True),
self._repos.quota_usage.count(system_id=system_id, near_critical=True),
self._repos.quota_usage.count(system_id=system_id, expiring_before=now),
)
total, by_window, near_warning_count, near_critical_count, expiring_count = await asyncio.gather(
self._repos.quota_usage.count(system_id=system_id),
self._repos.quota_usage.count_by_quota_window(system_id=system_id),
self._repos.quota_usage.count(system_id=system_id, near_warning=True),
self._repos.quota_usage.count(system_id=system_id, near_critical=True),
self._repos.quota_usage.count(system_id=system_id, expiring_before=now),
)
return QuotaStatsOutput(
total=total,