mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构
- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/ - 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层 - 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构 - 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations) - 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image - 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
15
_deprecated_py_src/services/user/__init__.py
Normal file
15
_deprecated_py_src/services/user/__init__.py
Normal file
@@ -0,0 +1,15 @@
|
||||
"""
|
||||
用户服务模块
|
||||
|
||||
包含用户管理、API Key 管理等功能。
|
||||
"""
|
||||
|
||||
from src.services.user.apikey import ApiKeyService
|
||||
from src.services.user.preference import PreferenceService
|
||||
from src.services.user.service import UserService
|
||||
|
||||
__all__ = [
|
||||
"UserService",
|
||||
"ApiKeyService",
|
||||
"PreferenceService",
|
||||
]
|
||||
307
_deprecated_py_src/services/user/apikey.py
Normal file
307
_deprecated_py_src/services/user/apikey.py
Normal file
@@ -0,0 +1,307 @@
|
||||
"""
|
||||
API密钥管理服务
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, Usage
|
||||
from src.services.user.bulk_cleanup import pre_clean_api_key
|
||||
|
||||
|
||||
class ApiKeyService:
|
||||
"""API密钥管理服务"""
|
||||
|
||||
@staticmethod
|
||||
def create_api_key(
|
||||
db: Session,
|
||||
user_id: str, # UUID
|
||||
name: str | None = None,
|
||||
allowed_providers: list[str] | None = None,
|
||||
allowed_api_formats: list[str] | None = None,
|
||||
allowed_models: list[str] | None = None,
|
||||
rate_limit: int | None = None,
|
||||
concurrent_limit: int = 5,
|
||||
expire_days: int | None = None,
|
||||
expires_at: datetime | None = None, # 直接传入过期时间,优先于 expire_days
|
||||
is_standalone: bool = False,
|
||||
auto_delete_on_expiry: bool = False,
|
||||
) -> tuple[ApiKey, str]:
|
||||
"""创建新的API密钥,返回密钥对象和明文密钥
|
||||
|
||||
Args:
|
||||
db: 数据库会话
|
||||
user_id: 用户ID
|
||||
name: 密钥名称
|
||||
allowed_providers: 允许的提供商列表
|
||||
allowed_api_formats: 允许的 API 格式列表
|
||||
allowed_models: 允许的模型列表
|
||||
rate_limit: 速率限制
|
||||
concurrent_limit: 并发限制
|
||||
expire_days: 过期天数,None = 永不过期
|
||||
expires_at: 直接指定过期时间,优先于 expire_days
|
||||
is_standalone: 是否为独立余额Key(仅管理员可创建)
|
||||
auto_delete_on_expiry: 过期后是否自动删除(True=物理删除,False=仅禁用)
|
||||
"""
|
||||
|
||||
# 生成密钥
|
||||
key = ApiKey.generate_key()
|
||||
key_hash = ApiKey.hash_key(key)
|
||||
key_encrypted = crypto_service.encrypt(key) # 加密存储密钥
|
||||
|
||||
# 计算过期时间:优先使用 expires_at,其次使用 expire_days
|
||||
final_expires_at = expires_at
|
||||
if final_expires_at is None and expire_days:
|
||||
final_expires_at = datetime.now(timezone.utc) + timedelta(days=expire_days)
|
||||
|
||||
normalized_rate_limit = rate_limit
|
||||
if not is_standalone and normalized_rate_limit is None:
|
||||
normalized_rate_limit = 0
|
||||
|
||||
api_key = ApiKey(
|
||||
user_id=user_id,
|
||||
key_hash=key_hash,
|
||||
key_encrypted=key_encrypted,
|
||||
name=name or f"API Key {datetime.now(timezone.utc).strftime('%Y%m%d%H%M%S')}",
|
||||
allowed_providers=allowed_providers,
|
||||
allowed_api_formats=allowed_api_formats,
|
||||
allowed_models=allowed_models,
|
||||
rate_limit=normalized_rate_limit,
|
||||
concurrent_limit=concurrent_limit,
|
||||
expires_at=final_expires_at,
|
||||
is_standalone=is_standalone,
|
||||
auto_delete_on_expiry=auto_delete_on_expiry,
|
||||
is_active=True,
|
||||
)
|
||||
|
||||
db.add(api_key)
|
||||
db.commit()
|
||||
db.refresh(api_key)
|
||||
|
||||
logger.info(
|
||||
f"创建API密钥: 用户ID {user_id}, 密钥名 {api_key.name}, " f"独立Key={is_standalone}"
|
||||
)
|
||||
return api_key, key # 返回密钥对象和明文密钥
|
||||
|
||||
@staticmethod
|
||||
def get_api_key(db: Session, key_id: str) -> ApiKey | None: # UUID
|
||||
"""获取API密钥"""
|
||||
return db.query(ApiKey).filter(ApiKey.id == key_id).first()
|
||||
|
||||
@staticmethod
|
||||
def get_api_key_by_key(db: Session, key: str) -> ApiKey | None:
|
||||
"""通过密钥字符串获取API密钥"""
|
||||
key_hash = ApiKey.hash_key(key)
|
||||
return db.query(ApiKey).filter(ApiKey.key_hash == key_hash).first()
|
||||
|
||||
@staticmethod
|
||||
def list_user_api_keys(
|
||||
db: Session, user_id: str, is_active: bool | None = None # UUID
|
||||
) -> list[ApiKey]:
|
||||
"""列出用户的所有API密钥(不包括独立Key)"""
|
||||
query = db.query(ApiKey).filter(
|
||||
ApiKey.user_id == user_id, ApiKey.is_standalone == False # 排除独立Key
|
||||
)
|
||||
|
||||
if is_active is not None:
|
||||
query = query.filter(ApiKey.is_active == is_active)
|
||||
|
||||
return query.order_by(ApiKey.created_at.desc()).all()
|
||||
|
||||
@staticmethod
|
||||
def list_standalone_api_keys(db: Session, is_active: bool | None = None) -> list[ApiKey]:
|
||||
"""列出所有独立余额Key(仅管理员可用)"""
|
||||
query = db.query(ApiKey).filter(ApiKey.is_standalone == True)
|
||||
|
||||
if is_active is not None:
|
||||
query = query.filter(ApiKey.is_active == is_active)
|
||||
|
||||
return query.order_by(ApiKey.created_at.desc()).all()
|
||||
|
||||
@staticmethod
|
||||
def update_api_key(db: Session, key_id: str, **kwargs: Any) -> ApiKey | None: # UUID
|
||||
"""更新API密钥"""
|
||||
api_key = db.query(ApiKey).filter(ApiKey.id == key_id).first()
|
||||
if not api_key:
|
||||
return None
|
||||
|
||||
# 可更新的字段
|
||||
updatable_fields = [
|
||||
"name",
|
||||
"allowed_providers",
|
||||
"allowed_api_formats",
|
||||
"allowed_models",
|
||||
"rate_limit",
|
||||
"concurrent_limit",
|
||||
"is_active",
|
||||
"expires_at",
|
||||
"auto_delete_on_expiry",
|
||||
]
|
||||
|
||||
# 允许显式设置为空数组/None 的字段(NULL=不限制,[]=全部禁用)
|
||||
nullable_list_fields = {"allowed_providers", "allowed_api_formats", "allowed_models"}
|
||||
|
||||
# 允许显式设置为 None 的字段(如 expires_at=None 表示永不过期;
|
||||
# standalone rate_limit=None 表示继承系统默认)
|
||||
nullable_fields = {"expires_at", "rate_limit"}
|
||||
|
||||
for field, value in kwargs.items():
|
||||
if field not in updatable_fields:
|
||||
continue
|
||||
# 对于 nullable_list_fields,保留 None/[] 的语义差异
|
||||
if field in nullable_list_fields:
|
||||
setattr(api_key, field, value)
|
||||
elif field in nullable_fields:
|
||||
if field == "rate_limit" and not api_key.is_standalone and value is None:
|
||||
setattr(api_key, field, 0)
|
||||
continue
|
||||
setattr(api_key, field, value)
|
||||
elif value is not None:
|
||||
setattr(api_key, field, value)
|
||||
|
||||
api_key.updated_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
db.refresh(api_key)
|
||||
|
||||
logger.debug(f"更新API密钥: ID {key_id}")
|
||||
return api_key
|
||||
|
||||
@staticmethod
|
||||
def delete_api_key(db: Session, key_id: str) -> bool: # UUID
|
||||
"""删除API密钥(禁用)"""
|
||||
api_key = db.query(ApiKey).filter(ApiKey.id == key_id).first()
|
||||
if not api_key:
|
||||
return False
|
||||
|
||||
api_key.is_active = False
|
||||
api_key.updated_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
|
||||
logger.info(f"删除API密钥: ID {key_id}")
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def cleanup_expired_keys(db: Session, auto_delete: bool = False) -> int:
|
||||
"""清理过期的API密钥
|
||||
|
||||
Args:
|
||||
db: 数据库会话
|
||||
auto_delete: 全局默认行为(True=物理删除,False=仅禁用)
|
||||
单个Key的 auto_delete_on_expiry 字段会覆盖此设置
|
||||
|
||||
Returns:
|
||||
int: 清理的密钥数量
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
expired_keys = (
|
||||
db.query(ApiKey)
|
||||
.filter(ApiKey.expires_at <= now, ApiKey.is_active == True) # 只处理仍然活跃的
|
||||
.all()
|
||||
)
|
||||
|
||||
count = 0
|
||||
for api_key in expired_keys:
|
||||
# 优先使用Key自身的auto_delete_on_expiry设置,否则使用全局设置
|
||||
should_delete = (
|
||||
api_key.auto_delete_on_expiry
|
||||
if api_key.auto_delete_on_expiry is not None
|
||||
else auto_delete
|
||||
)
|
||||
|
||||
if should_delete:
|
||||
# 物理删除(Usage / RequestCandidate / VideoTask 等记录保留)
|
||||
pre_clean_api_key(db, api_key.id)
|
||||
db.delete(api_key)
|
||||
logger.info(
|
||||
f"删除过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
|
||||
f"过期时间 {api_key.expires_at}"
|
||||
)
|
||||
else:
|
||||
# 仅禁用
|
||||
api_key.is_active = False
|
||||
api_key.updated_at = now
|
||||
logger.info(
|
||||
f"禁用过期API密钥: ID {api_key.id}, 名称 {api_key.name}, "
|
||||
f"过期时间 {api_key.expires_at}"
|
||||
)
|
||||
count += 1
|
||||
|
||||
if count > 0:
|
||||
db.commit()
|
||||
|
||||
return count
|
||||
|
||||
@staticmethod
|
||||
def get_api_key_stats(
|
||||
db: Session,
|
||||
key_id: str, # UUID
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""获取API密钥使用统计"""
|
||||
|
||||
api_key = db.query(ApiKey).filter(ApiKey.id == key_id).first()
|
||||
if not api_key:
|
||||
return {}
|
||||
|
||||
query = db.query(Usage).filter(Usage.api_key_id == key_id)
|
||||
|
||||
if start_date:
|
||||
query = query.filter(Usage.created_at >= start_date)
|
||||
if end_date:
|
||||
query = query.filter(Usage.created_at <= end_date)
|
||||
|
||||
# 统计数据
|
||||
stats = db.query(
|
||||
func.count(Usage.id).label("requests"),
|
||||
func.sum(Usage.total_tokens).label("tokens"),
|
||||
func.sum(Usage.total_cost_usd).label("cost_usd"),
|
||||
func.avg(Usage.response_time_ms).label("avg_response_time"),
|
||||
).filter(Usage.api_key_id == key_id)
|
||||
|
||||
if start_date:
|
||||
stats = stats.filter(Usage.created_at >= start_date)
|
||||
if end_date:
|
||||
stats = stats.filter(Usage.created_at <= end_date)
|
||||
|
||||
result = stats.first()
|
||||
|
||||
# 按天统计
|
||||
daily_stats = db.query(
|
||||
func.date(Usage.created_at).label("date"),
|
||||
func.count(Usage.id).label("requests"),
|
||||
func.sum(Usage.total_tokens).label("tokens"),
|
||||
func.sum(Usage.total_cost_usd).label("cost_usd"),
|
||||
).filter(Usage.api_key_id == key_id)
|
||||
|
||||
if start_date:
|
||||
daily_stats = daily_stats.filter(Usage.created_at >= start_date)
|
||||
if end_date:
|
||||
daily_stats = daily_stats.filter(Usage.created_at <= end_date)
|
||||
|
||||
daily_stats = daily_stats.group_by(func.date(Usage.created_at)).all()
|
||||
|
||||
return {
|
||||
"key_id": key_id,
|
||||
"key_name": api_key.name,
|
||||
"total_requests": result.requests or 0,
|
||||
"total_tokens": result.tokens or 0,
|
||||
"total_cost_usd": float(result.cost_usd or 0),
|
||||
"avg_response_time_ms": float(result.avg_response_time or 0),
|
||||
"daily_stats": [
|
||||
{
|
||||
"date": stat.date.isoformat() if stat.date else None,
|
||||
"requests": stat.requests,
|
||||
"tokens": stat.tokens,
|
||||
"cost_usd": float(stat.cost_usd),
|
||||
}
|
||||
for stat in daily_stats
|
||||
],
|
||||
}
|
||||
111
_deprecated_py_src/services/user/bulk_cleanup.py
Normal file
111
_deprecated_py_src/services/user/bulk_cleanup.py
Normal file
@@ -0,0 +1,111 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.models.database import RequestCandidate, Usage
|
||||
|
||||
_POSTGRES_BATCH_SIZE = 2000
|
||||
_SQLITE_BATCH_SIZE = 900
|
||||
|
||||
|
||||
def _resolve_batch_size(db: Session) -> int:
|
||||
try:
|
||||
bind = db.get_bind()
|
||||
dialect_name = str(getattr(getattr(bind, "dialect", None), "name", "") or "").lower()
|
||||
except Exception:
|
||||
dialect_name = ""
|
||||
|
||||
if dialect_name == "sqlite":
|
||||
return _SQLITE_BATCH_SIZE
|
||||
return _POSTGRES_BATCH_SIZE
|
||||
|
||||
|
||||
def batch_nullify_fk(
|
||||
db: Session,
|
||||
model: type[Any],
|
||||
column_name: str,
|
||||
entity_id: str | None,
|
||||
) -> int:
|
||||
"""分批将大表外键置空,避免单个长事务阻塞删除流程。"""
|
||||
if not entity_id:
|
||||
return 0
|
||||
|
||||
column = getattr(model, column_name)
|
||||
primary_key_column = next(iter(model.__table__.primary_key.columns))
|
||||
batch_size = _resolve_batch_size(db)
|
||||
total_updated = 0
|
||||
batch_index = 0
|
||||
started_at = time.monotonic()
|
||||
|
||||
while True:
|
||||
batch_ids = [
|
||||
row[0]
|
||||
for row in db.query(primary_key_column)
|
||||
.filter(column == entity_id)
|
||||
.limit(batch_size)
|
||||
.all()
|
||||
]
|
||||
if not batch_ids:
|
||||
break
|
||||
|
||||
batch_index += 1
|
||||
batch_started_at = time.monotonic()
|
||||
updated = int(
|
||||
db.query(model)
|
||||
.filter(primary_key_column.in_(batch_ids))
|
||||
.update({column: None}, synchronize_session=False)
|
||||
or 0
|
||||
)
|
||||
db.commit()
|
||||
|
||||
total_updated += updated
|
||||
elapsed_ms = int((time.monotonic() - batch_started_at) * 1000)
|
||||
logger.info(
|
||||
"批量清理 {}.{}: batch={}, updated={}, entity_id={}, elapsed_ms={}",
|
||||
model.__tablename__,
|
||||
column_name,
|
||||
batch_index,
|
||||
updated,
|
||||
entity_id,
|
||||
elapsed_ms,
|
||||
)
|
||||
|
||||
if len(batch_ids) < batch_size:
|
||||
break
|
||||
|
||||
if total_updated > 0:
|
||||
total_elapsed_ms = int((time.monotonic() - started_at) * 1000)
|
||||
logger.info(
|
||||
"批量清理完成 {}.{}: total_updated={}, entity_id={}, elapsed_ms={}",
|
||||
model.__tablename__,
|
||||
column_name,
|
||||
total_updated,
|
||||
entity_id,
|
||||
total_elapsed_ms,
|
||||
)
|
||||
|
||||
return total_updated
|
||||
|
||||
|
||||
def pre_clean_api_key(db: Session, api_key_id: str | None) -> int:
|
||||
"""预清理 API Key 在大表中的外键引用,减少后续删除锁竞争。"""
|
||||
if not api_key_id:
|
||||
return 0
|
||||
|
||||
usage_rows = batch_nullify_fk(db, Usage, "api_key_id", api_key_id)
|
||||
candidate_rows = batch_nullify_fk(db, RequestCandidate, "api_key_id", api_key_id)
|
||||
total_rows = usage_rows + candidate_rows
|
||||
|
||||
if total_rows > 0:
|
||||
logger.info(
|
||||
"API Key 预清理完成: api_key_id={}, usage={}, request_candidates={}",
|
||||
api_key_id,
|
||||
usage_rows,
|
||||
candidate_rows,
|
||||
)
|
||||
|
||||
return total_rows
|
||||
139
_deprecated_py_src/services/user/preference.py
Normal file
139
_deprecated_py_src/services/user/preference.py
Normal file
@@ -0,0 +1,139 @@
|
||||
"""
|
||||
用户偏好设置服务
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.exceptions import NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.models.database import Provider, User, UserPreference
|
||||
from src.services.wallet import WalletService
|
||||
|
||||
|
||||
class PreferenceService:
|
||||
"""用户偏好设置服务"""
|
||||
|
||||
@staticmethod
|
||||
def get_or_create_preferences(db: Session, user_id: str) -> UserPreference: # UUID
|
||||
"""获取或创建用户偏好设置"""
|
||||
preferences = db.query(UserPreference).filter(UserPreference.user_id == user_id).first()
|
||||
|
||||
if not preferences:
|
||||
# 创建默认偏好设置
|
||||
preferences = UserPreference(
|
||||
user_id=user_id,
|
||||
theme="light",
|
||||
language="zh-CN",
|
||||
timezone="Asia/Shanghai",
|
||||
email_notifications=True,
|
||||
usage_alerts=True,
|
||||
announcement_notifications=True,
|
||||
)
|
||||
db.add(preferences)
|
||||
db.commit()
|
||||
db.refresh(preferences)
|
||||
logger.info(f"Created default preferences for user {user_id}")
|
||||
|
||||
return preferences
|
||||
|
||||
@staticmethod
|
||||
def update_preferences(
|
||||
db: Session,
|
||||
user_id: str, # UUID
|
||||
avatar_url: str | None = None,
|
||||
bio: str | None = None,
|
||||
default_provider_id: str | None = None, # UUID
|
||||
theme: str | None = None,
|
||||
language: str | None = None,
|
||||
timezone: str | None = None,
|
||||
email_notifications: bool | None = None,
|
||||
usage_alerts: bool | None = None,
|
||||
announcement_notifications: bool | None = None,
|
||||
) -> UserPreference:
|
||||
"""更新用户偏好设置"""
|
||||
preferences = PreferenceService.get_or_create_preferences(db, user_id)
|
||||
|
||||
# 更新提供的字段
|
||||
if avatar_url is not None:
|
||||
preferences.avatar_url = avatar_url
|
||||
if bio is not None:
|
||||
preferences.bio = bio
|
||||
if default_provider_id is not None:
|
||||
# 验证提供商是否存在且活跃
|
||||
provider = (
|
||||
db.query(Provider)
|
||||
.filter(Provider.id == default_provider_id, Provider.is_active == True)
|
||||
.first()
|
||||
)
|
||||
if not provider:
|
||||
raise NotFoundException("Provider not found or inactive")
|
||||
preferences.default_provider_id = default_provider_id
|
||||
if theme is not None:
|
||||
if theme not in ["light", "dark", "auto", "system"]:
|
||||
raise ValueError("Invalid theme. Must be 'light', 'dark', 'auto', or 'system'")
|
||||
preferences.theme = theme
|
||||
if language is not None:
|
||||
preferences.language = language
|
||||
if timezone is not None:
|
||||
preferences.timezone = timezone
|
||||
if email_notifications is not None:
|
||||
preferences.email_notifications = email_notifications
|
||||
if usage_alerts is not None:
|
||||
preferences.usage_alerts = usage_alerts
|
||||
if announcement_notifications is not None:
|
||||
preferences.announcement_notifications = announcement_notifications
|
||||
|
||||
db.commit()
|
||||
db.refresh(preferences)
|
||||
logger.info(f"Updated preferences for user {user_id}")
|
||||
|
||||
return preferences
|
||||
|
||||
@staticmethod
|
||||
def get_user_with_preferences(db: Session, user_id: str) -> dict: # UUID
|
||||
"""获取用户信息及其偏好设置"""
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
if not user:
|
||||
raise NotFoundException("User not found")
|
||||
|
||||
preferences = PreferenceService.get_or_create_preferences(db, user_id)
|
||||
wallet = WalletService.get_wallet(db, user_id=user.id)
|
||||
billing = WalletService.serialize_wallet_summary(wallet)
|
||||
|
||||
# 构建返回数据
|
||||
user_data = {
|
||||
"id": user.id,
|
||||
"email": user.email,
|
||||
"username": user.username,
|
||||
"role": user.role.value,
|
||||
"is_active": user.is_active,
|
||||
"created_at": user.created_at,
|
||||
"last_login_at": user.last_login_at,
|
||||
"auth_source": user.auth_source.value if user.auth_source else "local",
|
||||
"has_password": bool(user.password_hash),
|
||||
"preferences": {
|
||||
"avatar_url": preferences.avatar_url,
|
||||
"bio": preferences.bio,
|
||||
"default_provider": (
|
||||
preferences.default_provider.name if preferences.default_provider else None
|
||||
),
|
||||
"theme": preferences.theme,
|
||||
"language": preferences.language,
|
||||
"timezone": preferences.timezone,
|
||||
"notifications": {
|
||||
"email": preferences.email_notifications,
|
||||
"usage_alerts": preferences.usage_alerts,
|
||||
"announcements": preferences.announcement_notifications,
|
||||
},
|
||||
},
|
||||
"billing": billing,
|
||||
"stats": {
|
||||
"total_cost": billing["total_consumed"],
|
||||
"total_cost_all_time": billing["total_consumed"],
|
||||
"api_keys_count": len(user.api_keys),
|
||||
},
|
||||
}
|
||||
|
||||
return user_data
|
||||
479
_deprecated_py_src/services/user/service.py
Normal file
479
_deprecated_py_src/services/user/service.py
Normal file
@@ -0,0 +1,479 @@
|
||||
"""
|
||||
用户管理服务
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import and_, func, or_
|
||||
from sqlalchemy.orm import Session, contains_eager
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.core.validators import EmailValidator, PasswordValidator, UsernameValidator
|
||||
from src.models.database import ApiKey, GlobalModel, Model, Provider, Usage, User, UserRole
|
||||
from src.services.auth.session_service import SessionService
|
||||
from src.services.cache.user_cache import UserCacheService
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.user.bulk_cleanup import batch_nullify_fk, pre_clean_api_key
|
||||
from src.utils.async_utils import safe_create_task
|
||||
from src.utils.transaction_manager import retry_on_database_error, transactional
|
||||
|
||||
|
||||
class UserService:
|
||||
"""用户管理服务"""
|
||||
|
||||
@staticmethod
|
||||
@transactional()
|
||||
@retry_on_database_error(max_retries=3)
|
||||
def create_user(
|
||||
db: Session,
|
||||
email: str | None,
|
||||
username: str,
|
||||
password: str,
|
||||
role: UserRole = UserRole.USER,
|
||||
initial_gift_usd: float | None = 10.0,
|
||||
unlimited: bool = False,
|
||||
email_verified: bool = False,
|
||||
allowed_providers: list[str] | None = None,
|
||||
allowed_api_formats: list[str] | None = None,
|
||||
allowed_models: list[str] | None = None,
|
||||
rate_limit: int | None = None,
|
||||
) -> User:
|
||||
"""创建新用户。"""
|
||||
|
||||
# 验证邮箱格式(仅当提供邮箱时)
|
||||
if email is not None:
|
||||
valid, error_msg = EmailValidator.validate(email)
|
||||
if not valid:
|
||||
raise ValueError(error_msg)
|
||||
# 检查邮箱是否已存在
|
||||
if db.query(User).filter(User.email == email).first():
|
||||
raise ValueError(f"邮箱已存在: {email}")
|
||||
|
||||
# 验证用户名格式
|
||||
valid, error_msg = UsernameValidator.validate(username)
|
||||
if not valid:
|
||||
raise ValueError(error_msg)
|
||||
|
||||
# 验证密码复杂度
|
||||
policy_level = SystemConfigService.get_password_policy_level(db)
|
||||
valid, error_msg = PasswordValidator.validate(password, policy=policy_level)
|
||||
if not valid:
|
||||
raise ValueError(error_msg)
|
||||
|
||||
# 检查用户名是否已存在
|
||||
if db.query(User).filter(User.username == username).first():
|
||||
raise ValueError(f"用户名已存在: {username}")
|
||||
|
||||
user = User(
|
||||
email=email,
|
||||
email_verified=email_verified if email else False,
|
||||
username=username,
|
||||
role=role,
|
||||
is_active=True,
|
||||
allowed_providers=allowed_providers,
|
||||
allowed_api_formats=allowed_api_formats,
|
||||
allowed_models=allowed_models,
|
||||
rate_limit=rate_limit,
|
||||
)
|
||||
user.set_password(password)
|
||||
|
||||
db.add(user)
|
||||
db.flush()
|
||||
|
||||
from src.services.wallet import WalletService
|
||||
|
||||
WalletService.initialize_user_wallet(
|
||||
db,
|
||||
user=user,
|
||||
initial_gift_usd=initial_gift_usd,
|
||||
unlimited=unlimited,
|
||||
description="用户初始赠款",
|
||||
)
|
||||
|
||||
db.commit() # 立即提交事务,释放数据库锁
|
||||
db.refresh(user)
|
||||
|
||||
log_identifier = email if email else username
|
||||
logger.info(f"创建新用户: {log_identifier} (ID: {user.id}, 角色: {role.value})")
|
||||
return user
|
||||
|
||||
@staticmethod
|
||||
@transactional()
|
||||
def create_user_with_api_key(
|
||||
db: Session,
|
||||
email: str,
|
||||
username: str,
|
||||
password: str,
|
||||
api_key_name: str = "默认密钥",
|
||||
role: UserRole = UserRole.USER,
|
||||
initial_gift_usd: float | None = 10.0,
|
||||
unlimited: bool = False,
|
||||
concurrent_limit: int = 5,
|
||||
) -> tuple[User, ApiKey]:
|
||||
"""
|
||||
创建用户并同时创建API密钥(原子操作)
|
||||
|
||||
Args:
|
||||
db: 数据库会话
|
||||
email: 邮箱
|
||||
username: 用户名
|
||||
password: 密码
|
||||
api_key_name: API密钥名称
|
||||
role: 用户角色
|
||||
initial_gift_usd: 初始赠款(USD)
|
||||
unlimited: 是否无限制
|
||||
concurrent_limit: 并发限制
|
||||
|
||||
Returns:
|
||||
tuple[User, ApiKey]: 用户对象和API密钥对象
|
||||
|
||||
Raises:
|
||||
ValueError: 当验证失败时
|
||||
"""
|
||||
# 创建用户
|
||||
user = UserService.create_user(
|
||||
db=db,
|
||||
email=email,
|
||||
username=username,
|
||||
password=password,
|
||||
role=role,
|
||||
initial_gift_usd=initial_gift_usd,
|
||||
unlimited=unlimited,
|
||||
)
|
||||
|
||||
# 导入API密钥服务(避免循环导入)
|
||||
from .apikey import ApiKeyService
|
||||
|
||||
# 创建API密钥(返回值是 (api_key, plain_key))
|
||||
api_key, plain_key = ApiKeyService.create_api_key(
|
||||
db=db, user_id=user.id, name=api_key_name, concurrent_limit=concurrent_limit
|
||||
)
|
||||
|
||||
logger.info(f"创建用户和API密钥完成: {email} (用户ID: {user.id}, 密钥ID: {api_key.id})")
|
||||
|
||||
# 返回用户对象、API Key对象和明文密钥
|
||||
return user, api_key, plain_key
|
||||
|
||||
@staticmethod
|
||||
def get_user(db: Session, user_id: str) -> User | None:
|
||||
"""获取用户"""
|
||||
import random
|
||||
import time
|
||||
|
||||
# 添加重试机制处理数据库并发问题
|
||||
max_retries = 3
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
return user
|
||||
except Exception as e:
|
||||
if attempt < max_retries - 1:
|
||||
# 添加随机延迟避免并发冲突
|
||||
time.sleep(random.uniform(0.01, 0.05))
|
||||
db.rollback() # 回滚事务
|
||||
continue
|
||||
else:
|
||||
raise e
|
||||
|
||||
@staticmethod
|
||||
def get_user_by_email(db: Session, email: str) -> User | None:
|
||||
"""通过邮箱获取用户"""
|
||||
return db.query(User).filter(User.email == email).first()
|
||||
|
||||
@staticmethod
|
||||
def list_users(
|
||||
db: Session,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
role: UserRole | None = None,
|
||||
is_active: bool | None = None,
|
||||
) -> list[User]:
|
||||
"""列出用户"""
|
||||
query = db.query(User)
|
||||
|
||||
if role:
|
||||
query = query.filter(User.role == role)
|
||||
if is_active is not None:
|
||||
query = query.filter(User.is_active == is_active)
|
||||
|
||||
return (
|
||||
query.order_by(User.created_at.desc(), User.id.desc()).offset(skip).limit(limit).all()
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@transactional()
|
||||
def update_user(db: Session, user_id: str, **kwargs: Any) -> User | None:
|
||||
"""更新用户信息"""
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
if not user:
|
||||
return None
|
||||
|
||||
# 可更新的字段
|
||||
updatable_fields = [
|
||||
"email",
|
||||
"username",
|
||||
"is_active",
|
||||
"role",
|
||||
# 访问限制字段
|
||||
"allowed_providers",
|
||||
"allowed_api_formats",
|
||||
"allowed_models",
|
||||
"rate_limit",
|
||||
]
|
||||
|
||||
# 允许设置为 None 的字段(表示无限制)
|
||||
nullable_fields = [
|
||||
"allowed_providers",
|
||||
"allowed_api_formats",
|
||||
"allowed_models",
|
||||
"rate_limit",
|
||||
]
|
||||
|
||||
for field, value in kwargs.items():
|
||||
if field not in updatable_fields:
|
||||
continue
|
||||
# nullable_fields 中的字段允许设置为 None
|
||||
if field in nullable_fields:
|
||||
setattr(user, field, value)
|
||||
elif value is not None:
|
||||
setattr(user, field, value)
|
||||
|
||||
# 如果提供了新密码
|
||||
if "password" in kwargs and kwargs["password"]:
|
||||
# 验证新密码复杂度
|
||||
policy_level = SystemConfigService.get_password_policy_level(db)
|
||||
valid, error_msg = PasswordValidator.validate(kwargs["password"], policy=policy_level)
|
||||
if not valid:
|
||||
raise ValueError(error_msg)
|
||||
user.set_password(kwargs["password"])
|
||||
SessionService.revoke_all_user_sessions(
|
||||
db,
|
||||
user_id=user.id,
|
||||
reason="admin_password_reset",
|
||||
)
|
||||
|
||||
user.updated_at = datetime.now(timezone.utc)
|
||||
db.commit() # 立即提交事务,释放数据库锁
|
||||
db.refresh(user)
|
||||
|
||||
# 清除用户缓存
|
||||
safe_create_task(UserCacheService.invalidate_user_cache(user.id, user.email))
|
||||
|
||||
logger.debug(f"更新用户信息: {user.email} (ID: {user_id})")
|
||||
return user
|
||||
|
||||
@staticmethod
|
||||
def delete_user(db: Session, user_id: str) -> bool:
|
||||
"""删除用户(硬删除)
|
||||
|
||||
删除流程:
|
||||
1. 检查未完结账务,阻止删除
|
||||
2. 预清理 Usage / RequestCandidate 等大表外键
|
||||
3. 手动删除 ORM cascade 冲突的子记录
|
||||
4. 删除用户记录
|
||||
5. 财务记录(Wallet/PaymentOrder/RefundRequest/WalletTransaction)和
|
||||
Usage 记录保留,外键 SET NULL,由自动清理策略统一回收
|
||||
"""
|
||||
from src.models.database import (
|
||||
AnnouncementRead,
|
||||
ApiKey,
|
||||
PaymentOrder,
|
||||
RefundRequest,
|
||||
RequestCandidate,
|
||||
UserPreference,
|
||||
Wallet,
|
||||
)
|
||||
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
if not user:
|
||||
return False
|
||||
|
||||
# 记录删除信息用于日志
|
||||
email = user.email
|
||||
|
||||
# 删除前阻断未完结账务,避免删除导致资金状态不一致。
|
||||
wallet_ids = [
|
||||
wallet_id
|
||||
for (wallet_id,) in (
|
||||
db.query(Wallet.id)
|
||||
.outerjoin(ApiKey, Wallet.api_key_id == ApiKey.id)
|
||||
.filter(or_(Wallet.user_id == user_id, ApiKey.user_id == user_id))
|
||||
.all()
|
||||
)
|
||||
]
|
||||
if wallet_ids:
|
||||
pending_refund_count = (
|
||||
db.query(RefundRequest)
|
||||
.filter(
|
||||
RefundRequest.wallet_id.in_(wallet_ids),
|
||||
RefundRequest.status.in_(["pending_approval", "approved", "processing"]),
|
||||
)
|
||||
.count()
|
||||
)
|
||||
if pending_refund_count > 0:
|
||||
raise ValueError("用户存在未完结退款,禁止删除")
|
||||
|
||||
pending_order_count = (
|
||||
db.query(PaymentOrder)
|
||||
.filter(
|
||||
PaymentOrder.wallet_id.in_(wallet_ids),
|
||||
PaymentOrder.status.in_(["pending", "paid"]),
|
||||
)
|
||||
.count()
|
||||
)
|
||||
if pending_order_count > 0:
|
||||
raise ValueError("用户存在未完结充值订单,禁止删除")
|
||||
|
||||
api_key_ids = [
|
||||
api_key_id
|
||||
for (api_key_id,) in db.query(ApiKey.id).filter(ApiKey.user_id == user_id).all()
|
||||
]
|
||||
api_key_count = len(api_key_ids)
|
||||
|
||||
# 注意:batch_nullify_fk 内部分批 commit,预清理部分不可回滚。
|
||||
# 这是预期行为:SET NULL 是幂等操作,即使后续步骤失败,
|
||||
# 已置空的外键不影响数据完整性,重新执行删除即可。
|
||||
try:
|
||||
for api_key_id in api_key_ids:
|
||||
pre_clean_api_key(db, api_key_id)
|
||||
|
||||
batch_nullify_fk(db, Usage, "user_id", user_id)
|
||||
batch_nullify_fk(db, RequestCandidate, "user_id", user_id)
|
||||
|
||||
db.query(UserPreference).filter(UserPreference.user_id == user_id).delete(
|
||||
synchronize_session=False
|
||||
)
|
||||
db.query(AnnouncementRead).filter(AnnouncementRead.user_id == user_id).delete(
|
||||
synchronize_session=False
|
||||
)
|
||||
|
||||
# 财务记录(Wallet/WalletTransaction/PaymentOrder/RefundRequest/PaymentCallback)
|
||||
# 和 Usage / RequestCandidate / VideoTask 记录全部保留,数据库外键 SET NULL 自动断开关联。
|
||||
db.query(ApiKey).filter(ApiKey.user_id == user_id).delete(synchronize_session=False)
|
||||
|
||||
# 现在删除用户(Usage, AuditLog, RequestAttempt 等会通过数据库 SET NULL 保留)
|
||||
db.delete(user)
|
||||
db.commit() # 立即提交事务,释放数据库锁
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
# 清除用户缓存
|
||||
safe_create_task(UserCacheService.invalidate_user_cache(user_id, email))
|
||||
|
||||
logger.info(f"删除用户: {email} (ID: {user_id}), 同时删除 {api_key_count} 个API密钥")
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def get_user_usage_stats(
|
||||
db: Session,
|
||||
user_id: str,
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""获取用户使用统计"""
|
||||
|
||||
query = db.query(Usage).filter(Usage.user_id == user_id)
|
||||
|
||||
if start_date:
|
||||
query = query.filter(Usage.created_at >= start_date)
|
||||
if end_date:
|
||||
query = query.filter(Usage.created_at <= end_date)
|
||||
|
||||
# 统计数据
|
||||
stats = db.query(
|
||||
func.count(Usage.id).label("total_requests"),
|
||||
func.sum(Usage.total_tokens).label("total_tokens"),
|
||||
func.sum(Usage.total_cost_usd).label("total_cost_usd"),
|
||||
func.avg(Usage.response_time_ms).label("avg_response_time"),
|
||||
).filter(Usage.user_id == user_id)
|
||||
|
||||
if start_date:
|
||||
stats = stats.filter(Usage.created_at >= start_date)
|
||||
if end_date:
|
||||
stats = stats.filter(Usage.created_at <= end_date)
|
||||
|
||||
result = stats.first()
|
||||
|
||||
# 按模型分组统计
|
||||
model_stats = db.query(
|
||||
Usage.model,
|
||||
func.count(Usage.id).label("requests"),
|
||||
func.sum(Usage.total_tokens).label("tokens"),
|
||||
func.sum(Usage.total_cost_usd).label("cost_usd"),
|
||||
).filter(Usage.user_id == user_id)
|
||||
|
||||
if start_date:
|
||||
model_stats = model_stats.filter(Usage.created_at >= start_date)
|
||||
if end_date:
|
||||
model_stats = model_stats.filter(Usage.created_at <= end_date)
|
||||
|
||||
model_stats = model_stats.group_by(Usage.model).all()
|
||||
|
||||
return {
|
||||
"total_requests": result.total_requests or 0,
|
||||
"total_tokens": result.total_tokens or 0,
|
||||
"total_cost_usd": float(result.total_cost_usd or 0),
|
||||
"avg_response_time_ms": float(result.avg_response_time or 0),
|
||||
"by_model": [
|
||||
{
|
||||
"model": stat.model,
|
||||
"requests": stat.requests,
|
||||
"tokens": stat.tokens,
|
||||
"cost_usd": float(stat.cost_usd),
|
||||
}
|
||||
for stat in model_stats
|
||||
],
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def get_user_available_models(db: Session, user: User) -> list[Model]:
|
||||
"""获取用户可用的模型
|
||||
|
||||
通过 GlobalModel + Model 关联查询用户可用模型
|
||||
逻辑:使用 AccessRestrictions 统一处理 allowed_providers 和 allowed_models 限制
|
||||
"""
|
||||
from src.core.access_restrictions import AccessRestrictions
|
||||
|
||||
# 使用 AccessRestrictions 类来处理限制(与 /v1/models 逻辑一致)
|
||||
restrictions = AccessRestrictions.from_api_key_and_user(api_key=None, user=user)
|
||||
|
||||
# 获取所有活跃的 Provider ID
|
||||
all_active_provider_ids = [
|
||||
p.id for p in db.query(Provider.id).filter(Provider.is_active == True).all()
|
||||
]
|
||||
|
||||
if not all_active_provider_ids:
|
||||
return []
|
||||
|
||||
# 查询所有活跃的 Model(关联 GlobalModel,contains_eager 避免循环中懒加载)
|
||||
all_models = (
|
||||
db.query(Model)
|
||||
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
|
||||
.options(contains_eager(Model.global_model))
|
||||
.filter(
|
||||
and_(
|
||||
Model.provider_id.in_(all_active_provider_ids),
|
||||
Model.is_active == True,
|
||||
GlobalModel.is_active == True,
|
||||
)
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 应用访问限制过滤
|
||||
filtered_models = []
|
||||
for model in all_models:
|
||||
model_name = (
|
||||
model.global_model.name if model.global_model else model.provider_model_name
|
||||
)
|
||||
# 使用 AccessRestrictions.is_model_allowed 检查模型是否可访问
|
||||
if restrictions.is_model_allowed(model_name, model.provider_id):
|
||||
filtered_models.append(model)
|
||||
|
||||
logger.debug(f"用户 {user.email} 可用模型: {len(filtered_models)} 个")
|
||||
|
||||
return filtered_models
|
||||
Reference in New Issue
Block a user