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:
fawney19
2026-04-03 16:26:16 +08:00
parent 8f26e1a31f
commit 1d9c77522a
868 changed files with 1735 additions and 2433 deletions

View File

@@ -0,0 +1,26 @@
"""
系统服务模块
包含系统配置、审计日志、公告等功能。
"""
from src.services.system.announcement import AnnouncementService
from src.services.system.audit import AuditService
from src.services.system.config import SystemConfigService
from src.services.system.maintenance_scheduler import (
MaintenanceScheduler,
get_maintenance_scheduler,
)
from src.services.system.scheduler import TaskScheduler, get_scheduler
from src.services.system.sync_stats import SyncStatsService
__all__ = [
"SystemConfigService",
"AuditService",
"AnnouncementService",
"MaintenanceScheduler",
"get_maintenance_scheduler",
"SyncStatsService",
"TaskScheduler",
"get_scheduler",
]

View File

@@ -0,0 +1,241 @@
"""
公告系统服务
"""
from __future__ import annotations
from datetime import datetime, timezone
from sqlalchemy import func, or_
from sqlalchemy.orm import Session
from src.core.exceptions import ForbiddenException, NotFoundException
from src.core.logger import logger
from src.models.database import Announcement, AnnouncementRead, User, UserRole
class AnnouncementService:
"""公告系统服务"""
@staticmethod
def create_announcement(
db: Session,
author_id: str, # UUID
title: str,
content: str,
type: str = "info",
priority: int = 0,
is_pinned: bool = False,
start_time: datetime | None = None,
end_time: datetime | None = None,
) -> Announcement:
"""创建公告"""
# 验证作者是否为管理员
author = db.query(User).filter(User.id == author_id).first()
if not author or author.role != UserRole.ADMIN:
raise ForbiddenException("Only administrators can create announcements")
# 验证类型
if type not in ["info", "warning", "maintenance", "important"]:
raise ValueError("Invalid announcement type")
announcement = Announcement(
title=title,
content=content,
type=type,
priority=priority,
author_id=author_id,
is_pinned=is_pinned,
start_time=start_time,
end_time=end_time,
is_active=True,
)
db.add(announcement)
db.commit()
db.refresh(announcement)
logger.info(f"Created announcement: {announcement.id} - {title}")
return announcement
@staticmethod
def get_announcements(
db: Session,
user_id: str | None = None, # UUID
active_only: bool = True,
include_read_status: bool = False,
limit: int = 50,
offset: int = 0,
) -> dict:
"""获取公告列表"""
query = db.query(Announcement)
# 筛选条件
if active_only:
now = datetime.now(timezone.utc)
query = query.filter(
Announcement.is_active == True,
or_(Announcement.start_time == None, Announcement.start_time <= now),
or_(Announcement.end_time == None, Announcement.end_time >= now),
)
# 分页
total = int(query.with_entities(func.count(Announcement.id)).scalar() or 0)
# 排序:置顶优先,然后按优先级和创建时间
query = query.order_by(
Announcement.is_pinned.desc(),
Announcement.priority.desc(),
Announcement.created_at.desc(),
)
announcements = query.offset(offset).limit(limit).all()
# 获取已读状态
read_announcement_ids = set()
unread_count = 0
if user_id and include_read_status:
read_records = (
db.query(AnnouncementRead.announcement_id)
.filter(AnnouncementRead.user_id == user_id)
.all()
)
read_announcement_ids = {r[0] for r in read_records}
unread_count = total - len(read_announcement_ids)
# 构建返回数据
items = []
for announcement in announcements:
item = {
"id": announcement.id,
"title": announcement.title,
"content": announcement.content,
"type": announcement.type,
"priority": announcement.priority,
"is_pinned": announcement.is_pinned,
"is_active": announcement.is_active,
"author": {"id": announcement.author.id, "username": announcement.author.username},
"start_time": announcement.start_time,
"end_time": announcement.end_time,
"created_at": announcement.created_at,
"updated_at": announcement.updated_at,
}
if include_read_status and user_id:
item["is_read"] = announcement.id in read_announcement_ids
items.append(item)
result = {"items": items, "total": total}
if include_read_status and user_id:
result["unread_count"] = unread_count
return result
@staticmethod
def get_announcement(db: Session, announcement_id: str) -> Announcement: # UUID
"""获取单个公告"""
announcement = db.query(Announcement).filter(Announcement.id == announcement_id).first()
if not announcement:
raise NotFoundException("Announcement not found")
return announcement
@staticmethod
def update_announcement(
db: Session,
announcement_id: str, # UUID
user_id: str, # UUID
title: str | None = None,
content: str | None = None,
type: str | None = None,
priority: int | None = None,
is_active: bool | None = None,
is_pinned: bool | None = None,
start_time: datetime | None = None,
end_time: datetime | None = None,
) -> Announcement:
"""更新公告"""
# 验证用户是否为管理员
user = db.query(User).filter(User.id == user_id).first()
if not user or user.role != UserRole.ADMIN:
raise ForbiddenException("Only administrators can update announcements")
announcement = AnnouncementService.get_announcement(db, announcement_id)
# 更新提供的字段
if title is not None:
announcement.title = title
if content is not None:
announcement.content = content
if type is not None:
if type not in ["info", "warning", "maintenance", "important"]:
raise ValueError("Invalid announcement type")
announcement.type = type
if priority is not None:
announcement.priority = priority
if is_active is not None:
announcement.is_active = is_active
if is_pinned is not None:
announcement.is_pinned = is_pinned
if start_time is not None:
announcement.start_time = start_time
if end_time is not None:
announcement.end_time = end_time
db.commit()
db.refresh(announcement)
logger.info(f"Updated announcement: {announcement_id}")
return announcement
@staticmethod
def delete_announcement(db: Session, announcement_id: str, user_id: str) -> None: # UUID
"""删除公告"""
# 验证用户是否为管理员
user = db.query(User).filter(User.id == user_id).first()
if not user or user.role != UserRole.ADMIN:
raise ForbiddenException("Only administrators can delete announcements")
announcement = AnnouncementService.get_announcement(db, announcement_id)
db.delete(announcement)
db.commit()
logger.info(f"Deleted announcement: {announcement_id}")
@staticmethod
def mark_as_read(db: Session, announcement_id: str, user_id: str) -> None: # UUID
"""标记公告为已读"""
# 检查公告是否存在
announcement = AnnouncementService.get_announcement(db, announcement_id)
# 检查是否已经标记为已读
existing = (
db.query(AnnouncementRead)
.filter(
AnnouncementRead.user_id == user_id,
AnnouncementRead.announcement_id == announcement_id,
)
.first()
)
if not existing:
read_record = AnnouncementRead(user_id=user_id, announcement_id=announcement_id)
db.add(read_record)
db.commit()
logger.info(f"User {user_id} marked announcement {announcement_id} as read")
@staticmethod
def get_active_announcements(db: Session, user_id: str | None = None) -> dict: # UUID
"""获取当前有效的公告(首页展示用)"""
return AnnouncementService.get_announcements(
db=db,
user_id=user_id,
active_only=True,
include_read_status=True if user_id else False,
limit=10,
)

View File

@@ -0,0 +1,462 @@
"""
审计日志服务
记录所有重要操作和安全事件
"""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.database import get_db
from src.models.database import AuditEventType, AuditLog
# 审计模型已移至 src/models/database.py
class AuditService:
"""审计服务
事务策略:本服务不负责事务提交,由中间件统一管理。
所有方法只做 db.add/flush提交由请求结束时的中间件处理。
"""
@staticmethod
def log_event(
db: Session,
event_type: AuditEventType,
description: str,
user_id: str | None = None, # UUID
api_key_id: str | None = None, # UUID
ip_address: str | None = None,
user_agent: str | None = None,
request_id: str | None = None,
status_code: int | None = None,
error_message: str | None = None,
metadata: dict[str, Any] | None = None,
) -> AuditLog:
"""
记录审计事件
Args:
db: 数据库会话
event_type: 事件类型
description: 事件描述
user_id: 用户ID
api_key_id: API密钥ID
ip_address: IP地址
user_agent: 用户代理
request_id: 请求ID
status_code: 状态码
error_message: 错误消息
metadata: 额外元数据
Returns:
审计日志记录
Note:
不在此方法内提交事务,由调用方或中间件统一管理。
"""
audit_log = AuditLog(
event_type=event_type.value,
description=description,
user_id=user_id,
api_key_id=api_key_id,
ip_address=ip_address,
user_agent=user_agent,
request_id=request_id,
status_code=status_code,
error_message=error_message,
event_metadata=metadata,
)
db.add(audit_log)
# 使用 flush 使记录可见但不提交事务,事务由中间件统一管理
db.flush()
# 同时记录到系统日志
# 检查 metadata 中是否有 quiet_logging 标志(由高频轮询端点设置)
quiet_logging = metadata.get("quiet_logging", False) if metadata else False
if not quiet_logging:
log_message = (
f"AUDIT [{event_type.value}] - {description} | "
f"user_id={user_id}, ip={ip_address}"
)
if event_type in [
AuditEventType.UNAUTHORIZED_ACCESS,
AuditEventType.SUSPICIOUS_ACTIVITY,
]:
logger.warning(log_message)
elif event_type in [AuditEventType.LOGIN_FAILED, AuditEventType.REQUEST_FAILED]:
logger.info(log_message)
# request_success 已由 Pipeline 日志覆盖,不再重复输出到控制台
return audit_log
@staticmethod
def log_login_attempt(
db: Session,
email: str,
success: bool,
ip_address: str,
user_agent: str,
user_id: str | None = None, # UUID
error_reason: str | None = None,
) -> Any:
"""
记录登录尝试
Args:
db: 数据库会话
email: 登录邮箱
success: 是否成功
ip_address: IP地址
user_agent: 用户代理
user_id: 用户ID成功时
error_reason: 失败原因
"""
event_type = AuditEventType.LOGIN_SUCCESS if success else AuditEventType.LOGIN_FAILED
description = f"Login attempt for {email}"
if not success and error_reason:
description += f": {error_reason}"
AuditService.log_event(
db=db,
event_type=event_type,
description=description,
user_id=user_id,
ip_address=ip_address,
user_agent=user_agent,
metadata={"email": email},
)
@staticmethod
def log_api_request(
db: Session,
user_id: str, # UUID
api_key_id: str, # UUID
request_id: str,
model: str,
provider: str,
success: bool,
ip_address: str,
status_code: int,
error_message: str | None = None,
input_tokens: int | None = None,
output_tokens: int | None = None,
cost_usd: float | None = None,
) -> Any:
"""
记录API请求
Args:
db: 数据库会话
user_id: 用户ID
api_key_id: API密钥ID
request_id: 请求ID
model: 模型名称
provider: 提供商名称
success: 是否成功
ip_address: IP地址
status_code: 状态码
error_message: 错误消息
input_tokens: 输入tokens
output_tokens: 输出tokens
cost_usd: 成本(美元)
"""
event_type = AuditEventType.REQUEST_SUCCESS if success else AuditEventType.REQUEST_FAILED
description = f"API request to {provider}/{model}"
metadata = {"model": model, "provider": provider}
if input_tokens:
metadata["input_tokens"] = input_tokens
if output_tokens:
metadata["output_tokens"] = output_tokens
if cost_usd:
metadata["cost_usd"] = cost_usd
AuditService.log_event(
db=db,
event_type=event_type,
description=description,
user_id=user_id,
api_key_id=api_key_id,
request_id=request_id,
ip_address=ip_address,
status_code=status_code,
error_message=error_message,
metadata=metadata,
)
@staticmethod
def log_security_event(
db: Session,
event_type: AuditEventType,
description: str,
ip_address: str,
user_id: str | None = None, # UUID
severity: str = "medium",
details: dict[str, Any] | None = None,
) -> Any:
"""
记录安全事件
Args:
db: 数据库会话
event_type: 事件类型
description: 事件描述
ip_address: IP地址
user_id: 用户ID
severity: 严重程度 (low, medium, high, critical)
details: 详细信息
"""
event_metadata = {"severity": severity}
if details:
event_metadata.update(details)
AuditService.log_event(
db=db,
event_type=event_type,
description=description,
user_id=user_id,
ip_address=ip_address,
metadata=event_metadata,
)
# 对于高严重性事件,简化日志输出
if severity in ["high", "critical"]:
logger.error(f"安全告警 [{severity.upper()}]: {description}")
@staticmethod
def get_user_audit_logs(
db: Session,
user_id: str, # UUID
event_types: list[AuditEventType] | None = None,
limit: int = 100,
) -> list[AuditLog]:
"""
获取用户的审计日志
Args:
db: 数据库会话
user_id: 用户ID
event_types: 事件类型过滤
limit: 返回数量限制
Returns:
审计日志列表
"""
query = db.query(AuditLog).filter(AuditLog.user_id == user_id)
if event_types:
event_type_values = [et.value for et in event_types]
query = query.filter(AuditLog.event_type.in_(event_type_values))
return query.order_by(AuditLog.created_at.desc()).limit(limit).all()
@staticmethod
def get_suspicious_activities(db: Session, hours: int = 24, limit: int = 100) -> list[AuditLog]:
"""
获取可疑活动
Args:
db: 数据库会话
hours: 时间范围(小时)
limit: 返回数量限制
Returns:
可疑活动列表
"""
cutoff_time = datetime.now(timezone.utc) - timedelta(hours=hours)
suspicious_types = [
AuditEventType.SUSPICIOUS_ACTIVITY.value,
AuditEventType.UNAUTHORIZED_ACCESS.value,
AuditEventType.LOGIN_FAILED.value,
AuditEventType.REQUEST_RATE_LIMITED.value,
]
return (
db.query(AuditLog)
.filter(AuditLog.event_type.in_(suspicious_types), AuditLog.created_at >= cutoff_time)
.order_by(AuditLog.created_at.desc())
.limit(limit)
.all()
)
@staticmethod
def analyze_user_behavior(db: Session, user_id: str, days: int = 30) -> dict[str, Any]: # UUID
"""
分析用户行为
Args:
db: 数据库会话
user_id: 用户ID
days: 分析天数
Returns:
行为分析结果
"""
from sqlalchemy import func
cutoff_time = datetime.now(timezone.utc) - timedelta(days=days)
# 统计各种事件类型
event_counts = (
db.query(AuditLog.event_type, func.count(AuditLog.id).label("count"))
.filter(AuditLog.user_id == user_id, AuditLog.created_at >= cutoff_time)
.group_by(AuditLog.event_type)
.all()
)
# 统计失败请求
failed_requests = (
db.query(func.count(AuditLog.id))
.filter(
AuditLog.user_id == user_id,
AuditLog.event_type == AuditEventType.REQUEST_FAILED.value,
AuditLog.created_at >= cutoff_time,
)
.scalar()
)
# 统计成功请求
success_requests = (
db.query(func.count(AuditLog.id))
.filter(
AuditLog.user_id == user_id,
AuditLog.event_type == AuditEventType.REQUEST_SUCCESS.value,
AuditLog.created_at >= cutoff_time,
)
.scalar()
)
# 获取最近的可疑活动
recent_suspicious = int(
db.query(func.count(AuditLog.id))
.filter(
AuditLog.user_id == user_id,
AuditLog.event_type.in_(
[
AuditEventType.SUSPICIOUS_ACTIVITY.value,
AuditEventType.UNAUTHORIZED_ACCESS.value,
]
),
AuditLog.created_at >= cutoff_time,
)
.scalar()
or 0
)
return {
"user_id": user_id,
"period_days": days,
"event_counts": {event: count for event, count in event_counts},
"failed_requests": failed_requests or 0,
"success_requests": success_requests or 0,
"success_rate": (
success_requests / (success_requests + failed_requests)
if (success_requests + failed_requests) > 0
else 0
),
"suspicious_activities": recent_suspicious,
"analysis_time": datetime.now(timezone.utc).isoformat(),
}
@staticmethod
def log_event_auto(
event_type: AuditEventType,
description: str,
user_id: str | None = None,
api_key_id: str | None = None,
ip_address: str | None = None,
user_agent: str | None = None,
request_id: str | None = None,
status_code: int | None = None,
error_message: str | None = None,
event_metadata: dict[str, Any] | None = None,
db: Session | None = None,
) -> AuditLog | None:
"""
自动管理数据库会话的审计日志记录方法
适用于中间件等无法直接获取数据库会话的场景
Args:
event_type: 事件类型
description: 事件描述
user_id: 用户ID
api_key_id: API密钥ID
ip_address: IP地址
user_agent: 用户代理
request_id: 请求ID
status_code: 状态码
error_message: 错误消息
event_metadata: 额外元数据
db: 数据库会话(可选,如不提供则自动创建)
Returns:
审计日志记录
"""
# 如果提供了数据库会话,使用它(不自动提交)
if db is not None:
try:
audit_log = AuditService.log_event(
db=db,
event_type=event_type,
description=description,
user_id=user_id,
api_key_id=api_key_id,
ip_address=ip_address,
user_agent=user_agent,
request_id=request_id,
status_code=status_code,
error_message=error_message,
metadata=event_metadata,
)
# 注意:不在这里提交,让调用方决定何时提交
return audit_log
except Exception as e:
logger.error(f"Failed to log audit event: {e}")
return None
# 如果没有提供会话,自动创建并管理
db_session = None
try:
db_session = next(get_db())
audit_log = AuditService.log_event(
db=db_session,
event_type=event_type,
description=description,
user_id=user_id,
api_key_id=api_key_id,
ip_address=ip_address,
user_agent=user_agent,
request_id=request_id,
status_code=status_code,
error_message=error_message,
metadata=event_metadata,
)
db_session.commit()
return audit_log
except Exception as e:
logger.error(f"Failed to log audit event with auto session: {e}")
if db_session is not None:
db_session.rollback()
return None
finally:
if db_session is not None:
db_session.close()
# 全局审计服务实例
audit_service = AuditService()

View File

@@ -0,0 +1,616 @@
"""
系统配置服务
"""
from __future__ import annotations
import json
import time
from enum import Enum
from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.core.validators import PasswordPolicyLevel
from src.models.database import Provider, SystemConfig
REQUEST_RECORD_LEVEL_KEY = "request_record_level"
_LEGACY_REQUEST_LOG_LEVEL_KEY = "request_log_level"
class RequestRecordLevel(str, Enum):
"""请求记录级别(控制请求/响应详情入库)"""
BASIC = "basic" # 仅记录基本信息tokens、成本等
HEADERS = "headers" # 记录基本信息+请求/响应头(敏感信息会脱敏)
FULL = "full" # 记录完整请求和响应包含body敏感信息会脱敏
# 进程内缓存 TTL- 系统配置变化不频繁,使用较长的 TTL
_CONFIG_CACHE_TTL = 60 # 1 分钟
# 调度相关配置使用更短的缓存 TTL确保多 Worker 部署时快速收敛
# 当管理员切换调度模式时,其他 Worker 最多延迟 5 秒即可感知变更
_SCHEDULING_CONFIG_CACHE_TTL = 5 # 5 秒
# 需要跨 Worker 快速同步的调度相关配置 key
SCHEDULING_CONFIG_KEYS = frozenset({"scheduling_mode", "provider_priority_mode"})
# 进程内缓存存储: {key: (value, expire_time)}
_config_cache: dict[str, tuple[Any, float]] = {}
def _get_cached_config(key: str) -> tuple[bool, Any]:
"""从进程内缓存获取配置值
Returns:
(hit, value): hit=True 表示缓存命中value 为缓存的值
"""
if key in _config_cache:
value, expire_time = _config_cache[key]
if time.time() < expire_time:
return True, value
# 缓存过期,安全删除(避免并发时 KeyError
_config_cache.pop(key, None)
return False, None
def _set_cached_config(key: str, value: Any) -> None:
"""设置进程内缓存
调度相关配置使用更短的 TTL5秒确保多 Worker 部署时快速收敛。
"""
ttl = _SCHEDULING_CONFIG_CACHE_TTL if key in SCHEDULING_CONFIG_KEYS else _CONFIG_CACHE_TTL
_config_cache[key] = (value, time.time() + ttl)
def invalidate_config_cache(key: str | None = None) -> None:
"""清除配置缓存
Args:
key: 配置键,如果为 None 则清除所有缓存
"""
global _config_cache
if key is None:
_config_cache = {}
logger.debug("已清除所有系统配置缓存")
else:
# 使用 pop 安全删除,避免并发时 KeyError
if _config_cache.pop(key, None) is not None:
logger.debug(f"已清除系统配置缓存: {key}")
class SystemConfigService:
"""系统配置服务类"""
# 默认配置
DEFAULT_CONFIGS = {
# 站点信息
"site_name": {
"value": "Aether",
"description": "站点名称,显示在页面标题、导航栏、登录页面和邮件中",
},
"site_subtitle": {
"value": "AI Gateway",
"description": "站点副标题,显示在导航栏品牌名称下方",
},
"default_user_initial_gift_usd": {
"value": 10.0,
"description": "新用户默认初始赠款(美元)",
},
"password_policy_level": {
"value": PasswordPolicyLevel.WEAK.value,
"description": "密码策略等级weak(弱密码), medium(中等强度), strong(强密码)",
},
REQUEST_RECORD_LEVEL_KEY: {
"value": RequestRecordLevel.BASIC.value,
"description": "请求记录级别basic(基本信息), headers(含请求/响应头), full(完整请求/响应)",
},
"max_request_body_size": {
"value": 5242880, # 5MB
"description": "最大请求体记录大小字节超过此大小的请求体将被截断仅影响数据库记录不影响真实API请求",
},
"max_response_body_size": {
"value": 5242880, # 5MB
"description": "最大响应体记录大小字节超过此大小的响应体将被截断仅影响数据库记录不影响真实API响应",
},
"sensitive_headers": {
"value": ["authorization", "x-api-key", "api-key", "cookie", "set-cookie"],
"description": "敏感请求头列表,这些请求头会被脱敏处理",
},
# 分级清理策略
"detail_log_retention_days": {
"value": 7,
"description": "详细日志保留天数,超过此天数后压缩 request_body 和 response_body 到压缩字段",
},
"compressed_log_retention_days": {
"value": 30,
"description": "压缩记录保留天数,超过此天数后删除压缩的 body 字段保留headers和统计",
},
"header_retention_days": {
"value": 90,
"description": "请求头保留天数,超过此天数后清空 request_headers 和 response_headers 字段",
},
"log_retention_days": {
"value": 365,
"description": "请求记录保存天数,超过此天数后删除整条使用记录",
},
"enable_auto_cleanup": {
"value": True,
"description": "是否启用自动清理任务,每天凌晨执行分级清理",
},
"cleanup_batch_size": {
"value": 1000,
"description": "每批次清理的记录数,避免单次操作过大影响数据库性能",
},
"request_candidates_retention_days": {
"value": 30,
"description": "请求候选记录保留天数,超过此天数的 request_candidates 审计记录将被自动清理",
},
"request_candidates_cleanup_batch_size": {
"value": 5000,
"description": "请求候选记录每批次清理条数,使用独立批次控制大表删除压力",
},
"enable_provider_checkin": {
"value": True,
"description": "是否启用 Provider 自动签到任务",
},
"provider_checkin_time": {
"value": "01:05",
"description": "Provider 自动签到执行时间HH:MM 格式24小时制",
},
"provider_priority_mode": {
"value": "provider",
"description": "优先级策略provider(提供商优先模式) 或 global_key(全局Key优先模式)",
},
"scheduling_mode": {
"value": "cache_affinity",
"description": "调度模式fixed_order(固定顺序模式,严格按优先级顺序) 或 cache_affinity(缓存亲和模式优先使用已缓存的Provider)",
},
"auto_delete_expired_keys": {
"value": False,
"description": "是否自动删除过期的API KeyTrue=物理删除False=仅禁用),仅管理员可配置",
},
"email_suffix_mode": {
"value": "none",
"description": "邮箱后缀限制模式none(不限制), whitelist(白名单), blacklist(黑名单)",
},
"email_suffix_list": {
"value": [],
"description": "邮箱后缀列表,配合 email_suffix_mode 使用",
},
# 格式转换配置
"enable_format_conversion": {
"value": True,
"description": "格式转换总开关:开启时允许跨格式转换;关闭时禁止任何跨格式转换",
},
"keep_priority_on_conversion": {
"value": False,
"description": "格式转换时保持优先级:开启时需要转换的候选保持原优先级;关闭时降级到不需要转换的候选之后",
},
"audit_log_retention_days": {
"value": 30,
"description": "审计日志保留天数,超过此天数的审计日志将被自动清理",
},
"enable_db_maintenance": {
"value": True,
"description": "是否启用数据库表维护任务(定期 VACUUM ANALYZE 防止表和索引膨胀)",
},
# 系统代理
"system_proxy_node_id": {
"value": None,
"description": "系统默认代理节点 ID为空时直连。仅影响提供商出站请求大模型API/余额查询/OAuth不影响系统内部接口",
},
# SMTP 邮件配置
"smtp_host": {
"value": None,
"description": "SMTP 服务器地址",
},
"smtp_port": {
"value": 587,
"description": "SMTP 服务器端口",
},
"smtp_user": {
"value": None,
"description": "SMTP 用户名",
},
"smtp_password": {
"value": None,
"description": "SMTP 密码(加密存储)",
},
"smtp_use_tls": {
"value": True,
"description": "是否使用 STARTTLS",
},
"smtp_use_ssl": {
"value": False,
"description": "是否使用 SSL/TLS",
},
"smtp_from_email": {
"value": None,
"description": "发件人邮箱地址",
},
"smtp_from_name": {
"value": "Aether",
"description": "发件人名称",
},
# OAuth Token 刷新配置
"enable_oauth_token_refresh": {
"value": True,
"description": "是否启用 OAuth Token 自动刷新任务,主动刷新即将过期的 OAuth token",
},
}
@classmethod
def get_config(cls, db: Session, key: str, default: Any | None = None) -> Any | None:
"""获取系统配置值(带进程内缓存)"""
# Backward-compatible alias: request_log_level -> request_record_level
if key in {REQUEST_RECORD_LEVEL_KEY, _LEGACY_REQUEST_LOG_LEVEL_KEY}:
value = cls._get_request_record_level_raw(db)
if value is not None:
return value
if REQUEST_RECORD_LEVEL_KEY in cls.DEFAULT_CONFIGS:
value = cls.DEFAULT_CONFIGS[REQUEST_RECORD_LEVEL_KEY]["value"]
_set_cached_config(REQUEST_RECORD_LEVEL_KEY, value)
return value
return default
# 1. 检查进程内缓存
hit, cached_value = _get_cached_config(key)
if hit:
return cached_value
# 2. 查询数据库
config = db.query(SystemConfig).filter(SystemConfig.key == key).first()
if config:
_set_cached_config(key, config.value)
return config.value
# 3. 如果配置不存在,使用默认值
if key in cls.DEFAULT_CONFIGS:
value = cls.DEFAULT_CONFIGS[key]["value"]
_set_cached_config(key, value)
return value
return default
@classmethod
def get_configs(cls, db: Session, keys: list[str]) -> dict[str, Any]:
"""
批量获取系统配置值
Args:
db: 数据库会话
keys: 配置键列表
Returns:
配置键值字典
"""
result = {}
# 一次查询获取所有配置
configs = db.query(SystemConfig).filter(SystemConfig.key.in_(keys)).all()
config_map = {c.key: c.value for c in configs}
# 填充结果,不存在的使用默认值
for key in keys:
if key in config_map:
result[key] = config_map[key]
elif key in cls.DEFAULT_CONFIGS:
result[key] = cls.DEFAULT_CONFIGS[key]["value"]
else:
result[key] = None
return result
@staticmethod
def set_config(
db: Session, key: str, value: Any, description: str | None = None
) -> SystemConfig:
"""设置系统配置值"""
if key == "password_policy_level":
normalized = str(value).strip().lower() if value is not None else ""
value = (
PasswordPolicyLevel(normalized).value
if normalized
else PasswordPolicyLevel.WEAK.value
)
# Backward-compatible alias: request_log_level -> request_record_level
if key in {REQUEST_RECORD_LEVEL_KEY, _LEGACY_REQUEST_LOG_LEVEL_KEY}:
config = (
db.query(SystemConfig).filter(SystemConfig.key == REQUEST_RECORD_LEVEL_KEY).first()
)
legacy = (
db.query(SystemConfig)
.filter(SystemConfig.key == _LEGACY_REQUEST_LOG_LEVEL_KEY)
.first()
)
if config:
config.value = value
if description:
config.description = description
# 如果同时存在旧 key删除它避免混乱
if legacy:
db.delete(legacy)
elif legacy:
# 原地迁移旧 key -> 新 key
legacy.key = REQUEST_RECORD_LEVEL_KEY
legacy.value = value
if description:
legacy.description = description
config = legacy
else:
config = SystemConfig(
key=REQUEST_RECORD_LEVEL_KEY, value=value, description=description
)
db.add(config)
db.commit()
db.refresh(config)
invalidate_config_cache(REQUEST_RECORD_LEVEL_KEY)
invalidate_config_cache(_LEGACY_REQUEST_LOG_LEVEL_KEY)
return config
config = db.query(SystemConfig).filter(SystemConfig.key == key).first()
if config:
# 更新现有配置
config.value = value
if description:
config.description = description
else:
# 创建新配置
config = SystemConfig(key=key, value=value, description=description)
db.add(config)
db.commit()
db.refresh(config)
# 清除缓存
invalidate_config_cache(key)
return config
@staticmethod
def get_password_policy_level(db: Session) -> str:
"""获取密码策略等级,异常值自动回退为弱策略。"""
value = SystemConfigService.get_config(
db, "password_policy_level", PasswordPolicyLevel.WEAK.value
)
return (
PasswordPolicyLevel(value).value
if value in PasswordPolicyLevel._value2member_map_
else PasswordPolicyLevel.WEAK.value
)
@staticmethod
def get_default_provider(db: Session) -> str | None:
"""
获取系统默认提供商
优先级1. 管理员设置的默认提供商 2. 数据库中第一个可用提供商
"""
# 首先尝试获取管理员设置的默认提供商
default_provider = SystemConfigService.get_config(db, "default_provider")
if default_provider:
return default_provider
# 如果没有设置fallback到数据库中第一个可用提供商
first_provider = db.query(Provider).filter(Provider.is_active == True).first()
if first_provider:
return first_provider.name
return None
@staticmethod
def set_default_provider(db: Session, provider_name: str) -> SystemConfig:
"""设置系统默认提供商"""
return SystemConfigService.set_config(
db, "default_provider", provider_name, "系统默认提供商,当用户未设置个人提供商时使用"
)
# 敏感配置项,不返回实际值
SENSITIVE_KEYS = {"smtp_password"}
@classmethod
def get_all_configs(cls, db: Session) -> list:
"""获取所有系统配置"""
configs = db.query(SystemConfig).all()
by_key = {c.key: c for c in configs}
result = []
for config in configs:
# Hide legacy key in list; present as canonical key instead.
if config.key == _LEGACY_REQUEST_LOG_LEVEL_KEY:
if REQUEST_RECORD_LEVEL_KEY in by_key:
continue
# Expose as canonical key name
config_key = REQUEST_RECORD_LEVEL_KEY
else:
config_key = config.key
item = {
"key": config_key,
"description": config.description,
"updated_at": config.updated_at.isoformat(),
}
# 对敏感配置,只返回是否已设置的标志,不返回实际值
if config.key in cls.SENSITIVE_KEYS:
item["value"] = None
item["is_set"] = bool(config.value)
else:
item["value"] = config.value
result.append(item)
return result
@classmethod
def delete_config(cls, db: Session, key: str) -> bool:
"""删除系统配置"""
# Backward-compatible alias: request_log_level -> request_record_level
if key in {REQUEST_RECORD_LEVEL_KEY, _LEGACY_REQUEST_LOG_LEVEL_KEY}:
configs = (
db.query(SystemConfig)
.filter(
SystemConfig.key.in_([REQUEST_RECORD_LEVEL_KEY, _LEGACY_REQUEST_LOG_LEVEL_KEY])
)
.all()
)
if not configs:
return False
for c in configs:
db.delete(c)
db.commit()
invalidate_config_cache(REQUEST_RECORD_LEVEL_KEY)
invalidate_config_cache(_LEGACY_REQUEST_LOG_LEVEL_KEY)
return True
config = db.query(SystemConfig).filter(SystemConfig.key == key).first()
if config:
db.delete(config)
db.commit()
# 清除缓存
invalidate_config_cache(key)
return True
return False
@classmethod
def init_default_configs(cls, db: Session) -> None:
"""初始化默认配置"""
for key, default_config in cls.DEFAULT_CONFIGS.items():
if not db.query(SystemConfig).filter(SystemConfig.key == key).first():
config = SystemConfig(
key=key,
value=default_config["value"],
description=default_config["description"],
)
db.add(config)
db.commit()
logger.info("初始化默认系统配置完成")
@classmethod
def _get_request_record_level_raw(cls, db: Session) -> Any | None:
"""Raw value from DB/cache for request record level (supports legacy key)."""
hit, cached_value = _get_cached_config(REQUEST_RECORD_LEVEL_KEY)
if hit:
return cached_value
config = db.query(SystemConfig).filter(SystemConfig.key == REQUEST_RECORD_LEVEL_KEY).first()
if config:
_set_cached_config(REQUEST_RECORD_LEVEL_KEY, config.value)
return config.value
hit, cached_value = _get_cached_config(_LEGACY_REQUEST_LOG_LEVEL_KEY)
if hit:
_set_cached_config(REQUEST_RECORD_LEVEL_KEY, cached_value)
return cached_value
legacy = (
db.query(SystemConfig).filter(SystemConfig.key == _LEGACY_REQUEST_LOG_LEVEL_KEY).first()
)
if legacy:
_set_cached_config(_LEGACY_REQUEST_LOG_LEVEL_KEY, legacy.value)
_set_cached_config(REQUEST_RECORD_LEVEL_KEY, legacy.value)
return legacy.value
return None
@classmethod
def get_request_record_level(cls, db: Session) -> RequestRecordLevel:
"""获取请求记录级别(控制请求/响应详情入库)"""
level = cls.get_config(db, REQUEST_RECORD_LEVEL_KEY, RequestRecordLevel.BASIC.value)
if isinstance(level, str):
return RequestRecordLevel(level)
return level
@classmethod
def get_log_level(cls, db: Session) -> RequestRecordLevel:
"""Deprecated: use get_request_record_level."""
return cls.get_request_record_level(db)
@classmethod
def should_log_headers(cls, db: Session) -> bool:
"""是否应该记录请求头"""
level = cls.get_request_record_level(db)
return level in [RequestRecordLevel.HEADERS, RequestRecordLevel.FULL]
@classmethod
def should_log_body(cls, db: Session) -> bool:
"""是否应该记录请求体和响应体"""
level = cls.get_request_record_level(db)
return level == RequestRecordLevel.FULL
@classmethod
def should_mask_sensitive_data(cls, db: Session) -> bool:
"""是否应该脱敏敏感数据(始终脱敏)"""
_ = db # 保持接口一致性
return True
@classmethod
def get_sensitive_headers(cls, db: Session) -> list:
"""获取敏感请求头列表"""
return cls.get_config(db, "sensitive_headers", [])
@classmethod
def is_format_conversion_enabled(cls, db: Session) -> bool:
"""检查全局格式转换是否启用"""
return bool(cls.get_config(db, "enable_format_conversion", True))
@classmethod
def is_keep_priority_on_conversion(cls, db: Session) -> bool:
"""检查格式转换时是否保持优先级"""
return bool(cls.get_config(db, "keep_priority_on_conversion", False))
@classmethod
def mask_sensitive_headers(cls, db: Session, headers: dict[str, Any]) -> dict[str, Any]:
"""脱敏敏感请求头"""
if not cls.should_mask_sensitive_data(db):
return headers
sensitive_headers = cls.get_sensitive_headers(db)
sensitive_lower = {h.lower() for h in sensitive_headers if isinstance(h, str) and h}
masked_headers = {}
for key, value in headers.items():
if key.lower() in sensitive_lower:
# 保留前后各4个字符中间用星号替换
if len(str(value)) > 8:
masked_value = str(value)[:4] + "****" + str(value)[-4:]
else:
masked_value = "****"
masked_headers[key] = masked_value
else:
masked_headers[key] = value
return masked_headers
@classmethod
def truncate_body(cls, db: Session, body: Any, is_request: bool = True) -> Any:
"""截断过大的请求体或响应体"""
max_size_key = "max_request_body_size" if is_request else "max_response_body_size"
max_size = cls.get_config(db, max_size_key, 5242880) # 5MB
if not body:
return body
# 转换为字符串以计算大小
body_str = json.dumps(body) if isinstance(body, (dict, list)) else str(body)
if len(body_str) > max_size:
# 截断并添加提示
truncated_str = body_str[:max_size]
if isinstance(body, (dict, list)):
try:
# 尝试保持JSON格式
return {
"_truncated": True,
"_original_size": len(body_str),
"_content": truncated_str,
}
except:
pass
return truncated_str + f"\n... (truncated, original size: {len(body_str)} bytes)"
return body

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,303 @@
"""
统一定时任务调度器
使用 APScheduler 管理所有定时任务,支持时区配置。
所有定时任务使用应用时区APP_TIMEZONE配置执行时间
数据存储仍然使用 UTC。
"""
from __future__ import annotations
from datetime import datetime
from typing import Any, Callable
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from apscheduler.triggers.cron import CronTrigger
from apscheduler.triggers.date import DateTrigger
from apscheduler.triggers.interval import IntervalTrigger
from src.config import config
from src.core.logger import logger
# 统一从 config 读取,不再重复 os.getenv
APP_TIMEZONE = config.app_timezone
class TaskScheduler:
"""统一定时任务调度器"""
_instance: TaskScheduler | None = None
def __init__(self) -> None:
self.scheduler = AsyncIOScheduler(timezone=APP_TIMEZONE)
self._started = False
@classmethod
def get_instance(cls) -> TaskScheduler:
"""获取调度器单例"""
if cls._instance is None:
cls._instance = TaskScheduler()
return cls._instance
def add_cron_job(
self,
func: Callable[..., Any],
hour: int | str,
minute: int = 0,
day_of_week: str | int | None = None,
job_id: str | None = None,
name: str | None = None,
timezone: str | None = None,
**kwargs: Any,
) -> Any:
"""
添加 cron 定时任务
Args:
func: 要执行的函数
hour: 执行时间(小时),使用业务时区
minute: 执行时间(分钟)
day_of_week: 星期几执行(如 "sun", "mon" 或 0-6
job_id: 任务ID
name: 任务名称(用于日志)
**kwargs: 传递给任务函数的参数
"""
trigger_timezone = timezone or APP_TIMEZONE
trigger = CronTrigger(
day_of_week=day_of_week, hour=hour, minute=minute, timezone=trigger_timezone
)
job_id = job_id or func.__name__
display_name = name or job_id
self.scheduler.add_job(
func,
trigger,
id=job_id,
name=display_name,
replace_existing=True,
kwargs=kwargs,
)
logger.info(
f"已注册定时任务: {display_name}, "
f"执行时间: {hour}:{minute:02d} ({trigger_timezone})"
)
def add_interval_job(
self,
func: Callable[..., Any],
seconds: int | None = None,
minutes: int | None = None,
hours: int | None = None,
job_id: str | None = None,
name: str | None = None,
**kwargs: Any,
) -> Any:
"""
添加间隔执行任务
Args:
func: 要执行的函数
seconds: 间隔秒数
minutes: 间隔分钟数
hours: 间隔小时数
job_id: 任务ID
name: 任务名称
**kwargs: 传递给任务函数的参数
"""
# 构建 trigger 参数,过滤掉 None 值
trigger_kwargs = {}
if seconds is not None:
trigger_kwargs["seconds"] = seconds
if minutes is not None:
trigger_kwargs["minutes"] = minutes
if hours is not None:
trigger_kwargs["hours"] = hours
trigger = IntervalTrigger(**trigger_kwargs)
job_id = job_id or func.__name__
display_name = name or job_id
# 计算间隔描述
interval_parts = []
if hours:
interval_parts.append(f"{hours}小时")
if minutes:
interval_parts.append(f"{minutes}分钟")
if seconds:
interval_parts.append(f"{seconds}")
interval_desc = "".join(interval_parts) or "未知间隔"
self.scheduler.add_job(
func,
trigger,
id=job_id,
name=display_name,
replace_existing=True,
kwargs=kwargs,
)
logger.info(f"已注册间隔任务: {display_name}, 执行间隔: {interval_desc}")
def add_date_job(
self,
func: Callable[..., Any],
run_date: datetime,
job_id: str | None = None,
name: str | None = None,
**kwargs: Any,
) -> Any:
"""
添加一次性定时任务(在指定时间执行一次)
Args:
func: 要执行的函数
run_date: 执行时间datetime 对象)
job_id: 任务ID
name: 任务名称(用于日志)
**kwargs: 传递给任务函数的参数
"""
trigger = DateTrigger(run_date=run_date)
job_id = job_id or func.__name__
display_name = name or job_id
self.scheduler.add_job(
func,
trigger,
id=job_id,
name=display_name,
replace_existing=True,
kwargs=kwargs,
)
logger.info("已注册一次性任务: {}, 执行时间: {}", display_name, run_date.isoformat())
def start(self) -> Any:
"""启动调度器"""
if self._started:
logger.warning("调度器已在运行中")
return
self.scheduler.start()
self._started = True
logger.info(f"定时任务调度器已启动,应用时区: {APP_TIMEZONE}")
# 打印下次执行时间
self._log_next_run_times()
def stop(self) -> Any:
"""停止调度器"""
if not self._started:
return
self.scheduler.shutdown(wait=False)
self._started = False
logger.info("定时任务调度器已停止")
def _log_next_run_times(self) -> None:
"""记录所有任务的下次执行时间"""
jobs = self.scheduler.get_jobs()
if not jobs:
return
logger.info("已注册的定时任务:")
for job in jobs:
next_run = job.next_run_time
if next_run:
# 计算距离下次执行的时间
now = datetime.now(next_run.tzinfo)
delta = next_run - now
hours, remainder = divmod(int(delta.total_seconds()), 3600)
minutes = remainder // 60
logger.info(
f" - {job.name}: 下次执行 {next_run.strftime('%Y-%m-%d %H:%M')} "
f"({hours}小时{minutes}分钟后)"
)
def remove_job(self, job_id: str) -> None:
"""
移除指定的定时任务
Args:
job_id: 任务ID
"""
try:
self.scheduler.remove_job(job_id)
logger.info(f"已移除定时任务: {job_id}")
except Exception as e:
logger.warning(f"移除定时任务失败 {job_id}: {e}")
def reschedule_cron_job(
self,
job_id: str,
hour: int,
minute: int = 0,
) -> bool:
"""
重新调度 cron 定时任务的执行时间
Args:
job_id: 任务ID
hour: 新的执行时间(小时),使用业务时区
minute: 新的执行时间(分钟)
Returns:
是否成功重新调度
"""
try:
job = self.scheduler.get_job(job_id)
if not job:
logger.warning(f"任务不存在: {job_id}")
return False
trigger = CronTrigger(hour=hour, minute=minute, timezone=APP_TIMEZONE)
self.scheduler.reschedule_job(job_id, trigger=trigger)
logger.info(
f"已重新调度定时任务: {job.name}, "
f"新执行时间: {hour:02d}:{minute:02d} ({APP_TIMEZONE})"
)
return True
except Exception as e:
logger.exception(f"重新调度任务失败 {job_id}: {e}")
return False
def get_job_info(self, job_id: str) -> dict | None:
"""
获取任务信息
Args:
job_id: 任务ID
Returns:
任务信息字典,包含 name, next_run_time 等
"""
try:
job = self.scheduler.get_job(job_id)
if not job:
return None
next_run = job.next_run_time
return {
"id": job.id,
"name": job.name,
"next_run_time": next_run.isoformat() if next_run else None,
}
except Exception as e:
logger.warning(f"获取任务信息失败 {job_id}: {e}")
return None
@property
def is_running(self) -> bool:
"""调度器是否在运行"""
return self._started
# 便捷函数
def get_scheduler() -> TaskScheduler:
"""获取调度器单例"""
return TaskScheduler.get_instance()

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,175 @@
"""
API密钥统计同步服务
定期同步API密钥的统计数据确保与实际使用记录一致
"""
from __future__ import annotations
from typing import Any
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, Usage
class SyncStatsService:
"""API密钥统计同步服务"""
# 分页批量大小
BATCH_SIZE = 100
@staticmethod
def sync_api_key_stats(db: Session, api_key_id: str | None = None) -> dict: # UUID
"""
同步API密钥的统计数据
Args:
db: 数据库会话
api_key_id: 指定要同步的API密钥ID如果不指定则同步所有
Returns:
同步结果统计
"""
result = {"synced": 0, "updated": 0, "errors": 0}
try:
# 获取要同步的API密钥使用分页避免大数据量问题
if api_key_id:
single_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
api_keys = [single_key] if single_key else []
else:
# 分页处理,避免一次加载所有数据
offset = 0
api_keys: list[ApiKey] = []
while True:
batch = db.query(ApiKey).offset(offset).limit(SyncStatsService.BATCH_SIZE).all()
if not batch:
break
api_keys.extend(batch)
offset += SyncStatsService.BATCH_SIZE
# Pre-aggregate Usage stats in ONE query to avoid per-key N+1 scans.
# This is critical for large datasets (DB CPU killer otherwise).
usage_stats_map: dict[str, dict[str, Any]] = {}
if not api_key_id:
rows = (
db.query(
Usage.api_key_id,
func.count(Usage.id).label("requests"),
func.sum(Usage.total_cost_usd).label("cost"),
func.max(Usage.created_at).label("last_used"),
)
.filter(Usage.api_key_id.isnot(None))
.group_by(Usage.api_key_id)
.all()
)
usage_stats_map = {
str(r.api_key_id): {
"requests": int(r.requests or 0),
"cost": float(r.cost or 0),
"last_used": r.last_used,
}
for r in rows
if r.api_key_id is not None
}
for api_key in api_keys:
try:
if api_key_id:
# 单 key 路径:直接查(数据量小)
stats = (
db.query(
func.count(Usage.id).label("requests"),
func.sum(Usage.total_cost_usd).label("cost"),
func.max(Usage.created_at).label("last_used"),
)
.filter(Usage.api_key_id == api_key.id)
.first()
)
actual_requests = int(stats.requests or 0) if stats else 0
actual_cost = float(stats.cost or 0) if stats else 0.0
last_used_at = stats.last_used if stats else None
else:
# 批量路径:使用预聚合结果
s = usage_stats_map.get(str(api_key.id)) or {}
actual_requests = int(s.get("requests") or 0)
actual_cost = float(s.get("cost") or 0.0)
last_used_at = s.get("last_used")
# 检查是否需要更新
needs_update = False
if api_key.total_requests != actual_requests:
logger.info(
f"API密钥 {api_key.id} 请求数不一致: {api_key.total_requests} -> {actual_requests}"
)
api_key.total_requests = actual_requests
needs_update = True
if abs(float(api_key.total_cost_usd or 0) - actual_cost) > 0.0001:
logger.info(
f"API密钥 {api_key.id} 费用不一致: {api_key.total_cost_usd} -> {actual_cost}"
)
api_key.total_cost_usd = actual_cost
needs_update = True
if last_used_at and api_key.last_used_at != last_used_at:
api_key.last_used_at = last_used_at
needs_update = True
result["synced"] += 1
if needs_update:
result["updated"] += 1
logger.info(f"已更新API密钥 {api_key.id} 的统计数据")
except Exception as e:
logger.error(f"同步API密钥 {api_key.id} 统计时出错: {e}")
result["errors"] += 1
# 回滚当前失败的操作,继续处理其他密钥
try:
db.rollback()
except Exception:
pass
# 提交所有更改
db.commit()
logger.info(
f"同步完成: 处理 {result['synced']} 个密钥, 更新 {result['updated']} 个, 错误 {result['errors']}"
)
except Exception as e:
logger.error(f"同步统计数据时出错: {e}")
db.rollback()
raise
return result
@staticmethod
def get_api_key_real_stats(db: Session, api_key_id: str) -> dict: # UUID
"""
获取API密钥的实际统计数据直接从使用记录计算
Args:
db: 数据库会话
api_key_id: API密钥ID
Returns:
实际的统计数据
"""
# 计算实际的使用统计
stats = (
db.query(
func.count(Usage.id).label("requests"),
func.sum(Usage.total_cost_usd).label("cost"),
func.max(Usage.created_at).label("last_used"),
)
.filter(Usage.api_key_id == api_key_id)
.first()
)
return {
"total_requests": stats.requests or 0,
"total_cost_usd": float(stats.cost or 0),
"last_used_at": stats.last_used,
}

View File

@@ -0,0 +1,241 @@
"""Time range utilities for stats queries."""
from __future__ import annotations
from datetime import date, datetime, time, timedelta, timezone
from typing import Literal
from pydantic import BaseModel, model_validator
class TimeRangeParams(BaseModel):
"""
Time range parameters (local-date semantics).
Rules:
1) Inputs are user-local dates, backend converts to UTC datetime range.
2) Range is half-open: [start, end).
"""
start_date: date | None = None
end_date: date | None = None
preset: (
Literal[
"today",
"yesterday",
"last7days",
"last30days",
"last90days",
"this_week",
"last_week",
"this_month",
"last_month",
"this_year",
]
| None
) = None
granularity: Literal["hour", "day", "week", "month"] = "day"
timezone: str | None = None
tz_offset_minutes: int = 0
@model_validator(mode="after")
def validate_and_resolve(self) -> "TimeRangeParams":
"""Validate and resolve preset to concrete dates."""
if self.preset:
user_today = self._get_user_today()
match self.preset:
case "today":
self.start_date = self.end_date = user_today
case "yesterday":
self.start_date = self.end_date = user_today - timedelta(days=1)
case "last7days":
self.start_date = user_today - timedelta(days=6)
self.end_date = user_today
case "last30days":
self.start_date = user_today - timedelta(days=29)
self.end_date = user_today
case "last90days":
self.start_date = user_today - timedelta(days=89)
self.end_date = user_today
case "this_week":
self.start_date = user_today - timedelta(days=user_today.weekday())
self.end_date = user_today
case "last_week":
week_start = user_today - timedelta(days=user_today.weekday())
self.start_date = week_start - timedelta(days=7)
self.end_date = week_start - timedelta(days=1)
case "this_month":
self.start_date = user_today.replace(day=1)
self.end_date = user_today
case "last_month":
first_of_this_month = user_today.replace(day=1)
self.end_date = first_of_this_month - timedelta(days=1)
self.start_date = self.end_date.replace(day=1)
case "this_year":
self.start_date = user_today.replace(month=1, day=1)
self.end_date = user_today
if not self.preset and (self.start_date is None or self.end_date is None):
raise ValueError("Either preset or both start_date and end_date must be provided")
if self.start_date and self.end_date and self.start_date > self.end_date:
raise ValueError("start_date must be <= end_date")
if self.start_date and self.end_date:
max_days = 365
days = (self.end_date - self.start_date).days
if days > max_days:
raise ValueError(f"Query range cannot exceed {max_days} days")
if self.granularity == "hour":
if self.start_date != self.end_date:
raise ValueError("Hour granularity only supports single day query")
return self
def validate_for_time_series(self) -> "TimeRangeParams":
"""Extra validation for time series queries."""
if self.granularity == "hour" and self.start_date != self.end_date:
raise ValueError("Hour granularity only supports single day query")
if self.start_date and self.end_date:
days_inclusive = (self.end_date - self.start_date).days + 1
max_days_for_time_series = 90
if days_inclusive > max_days_for_time_series:
raise ValueError(
f"Time series query range cannot exceed {max_days_for_time_series} days "
f"(requested {days_inclusive} days). "
"For longer ranges, use aggregated statistics instead."
)
return self
def _get_user_today(self) -> date:
"""Get user-local 'today'."""
if self.timezone:
try:
from zoneinfo import ZoneInfo
user_tz = ZoneInfo(self.timezone)
return datetime.now(user_tz).date()
except Exception:
pass
user_now = datetime.now(timezone.utc) + timedelta(minutes=self.tz_offset_minutes)
return user_now.date()
def _get_tz_offset_for_date(self, local_date: date) -> timedelta:
"""Get timezone offset for a local date (DST-aware if timezone is provided)."""
if self.timezone:
try:
from zoneinfo import ZoneInfo
user_tz = ZoneInfo(self.timezone)
local_midnight = datetime.combine(local_date, time.min)
local_aware = local_midnight.replace(tzinfo=user_tz)
return local_aware.utcoffset() or timedelta(0)
except Exception:
pass
return timedelta(minutes=self.tz_offset_minutes)
def to_utc_datetime_range(self) -> tuple[datetime, datetime]:
"""Convert to UTC datetime range (half-open)."""
start_offset = self._get_tz_offset_for_date(self.start_date)
end_offset = self._get_tz_offset_for_date(self.end_date + timedelta(days=1))
local_start = datetime.combine(self.start_date, time.min)
local_end = datetime.combine(self.end_date + timedelta(days=1), time.min)
start_utc = (local_start - start_offset).replace(tzinfo=timezone.utc)
end_utc = (local_end - end_offset).replace(tzinfo=timezone.utc)
return start_utc, end_utc
def get_complete_utc_dates(
self,
) -> tuple[list[date], tuple[datetime, datetime] | None, tuple[datetime, datetime] | None]:
"""Split into complete UTC days + head/tail boundaries."""
start_utc, end_utc = self.to_utc_datetime_range()
if (
start_utc.hour == 0
and start_utc.minute == 0
and start_utc.second == 0
and start_utc.microsecond == 0
):
first_complete_date = start_utc.date()
head_boundary = None
else:
first_complete_date = start_utc.date() + timedelta(days=1)
head_boundary = (
start_utc,
datetime.combine(first_complete_date, time.min, tzinfo=timezone.utc),
)
if (
end_utc.hour == 0
and end_utc.minute == 0
and end_utc.second == 0
and end_utc.microsecond == 0
):
last_complete_date = end_utc.date() - timedelta(days=1)
tail_boundary = None
else:
last_complete_date = end_utc.date() - timedelta(days=1)
tail_start = datetime.combine(end_utc.date(), time.min, tzinfo=timezone.utc)
tail_boundary = (tail_start, end_utc)
complete_dates = []
if first_complete_date <= last_complete_date:
current = first_complete_date
while current <= last_complete_date:
complete_dates.append(current)
current += timedelta(days=1)
return complete_dates, head_boundary, tail_boundary
def get_local_day_hours(self) -> list[tuple[date, datetime, datetime]]:
"""Return local-day mapped UTC ranges for time-series."""
result = []
current_date = self.start_date
while current_date <= self.end_date:
offset = self._get_tz_offset_for_date(current_date)
local_start = datetime.combine(current_date, time.min)
local_end = datetime.combine(current_date + timedelta(days=1), time.min)
day_start_utc = (local_start - offset).replace(tzinfo=timezone.utc)
day_end_utc = (
local_end - self._get_tz_offset_for_date(current_date + timedelta(days=1))
).replace(tzinfo=timezone.utc)
result.append((current_date, day_start_utc, day_end_utc))
current_date += timedelta(days=1)
return result
def split_time_range_for_hourly(start_utc: datetime, end_utc: datetime) -> tuple[
tuple[datetime, datetime] | None,
list[datetime],
tuple[datetime, datetime] | None,
]:
"""Split into head fragment, complete hours, tail fragment."""
first_hour = start_utc.replace(minute=0, second=0, microsecond=0)
if start_utc > first_hour:
first_hour += timedelta(hours=1)
head_fragment = (start_utc, first_hour) if first_hour <= end_utc else None
else:
head_fragment = None
first_hour = start_utc
last_hour = end_utc.replace(minute=0, second=0, microsecond=0)
if end_utc > last_hour:
tail_fragment = (last_hour, end_utc) if last_hour >= first_hour else None
else:
tail_fragment = None
last_hour = end_utc
complete_hours = []
current = first_hour
while current < last_hour:
complete_hours.append(current)
current += timedelta(hours=1)
return head_fragment, complete_hours, tail_fragment