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:
26
_deprecated_py_src/services/system/__init__.py
Normal file
26
_deprecated_py_src/services/system/__init__.py
Normal 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",
|
||||
]
|
||||
241
_deprecated_py_src/services/system/announcement.py
Normal file
241
_deprecated_py_src/services/system/announcement.py
Normal 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,
|
||||
)
|
||||
462
_deprecated_py_src/services/system/audit.py
Normal file
462
_deprecated_py_src/services/system/audit.py
Normal 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()
|
||||
616
_deprecated_py_src/services/system/config.py
Normal file
616
_deprecated_py_src/services/system/config.py
Normal 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:
|
||||
"""设置进程内缓存
|
||||
|
||||
调度相关配置使用更短的 TTL(5秒),确保多 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 Key(True=物理删除,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
|
||||
1478
_deprecated_py_src/services/system/maintenance_scheduler.py
Normal file
1478
_deprecated_py_src/services/system/maintenance_scheduler.py
Normal file
File diff suppressed because it is too large
Load Diff
303
_deprecated_py_src/services/system/scheduler.py
Normal file
303
_deprecated_py_src/services/system/scheduler.py
Normal 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()
|
||||
2014
_deprecated_py_src/services/system/stats_aggregator.py
Normal file
2014
_deprecated_py_src/services/system/stats_aggregator.py
Normal file
File diff suppressed because it is too large
Load Diff
175
_deprecated_py_src/services/system/sync_stats.py
Normal file
175
_deprecated_py_src/services/system/sync_stats.py
Normal 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,
|
||||
}
|
||||
241
_deprecated_py_src/services/system/time_range.py
Normal file
241
_deprecated_py_src/services/system/time_range.py
Normal 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
|
||||
Reference in New Issue
Block a user