Files
Aether/_deprecated_py_src/services/request/candidate.py
fawney19 1d9c77522a 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)
2026-04-03 16:26:16 +08:00

404 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
请求候选记录服务 - 管理候选队列
"""
import uuid
from datetime import datetime, timezone
from sqlalchemy.orm import Session
from src.core.batch_committer import get_batch_committer
from src.core.logger import logger
from src.models.database import RequestCandidate
class RequestCandidateService:
"""请求候选记录服务"""
@staticmethod
def _persist_candidate_update(db: Session, *, immediate: bool) -> None:
if immediate:
db.commit()
return
db.flush()
get_batch_committer().mark_dirty(db)
@staticmethod
def create_candidate(
db: Session,
request_id: str,
candidate_index: int,
candidate_id: str | None = None,
retry_index: int = 0, # 新增:重试序号
user_id: str | None = None,
api_key_id: str | None = None,
username: str | None = None,
api_key_name: str | None = None,
provider_id: str | None = None,
endpoint_id: str | None = None,
key_id: str | None = None,
status: str = "available",
skip_reason: str | None = None,
is_cached: bool = False,
extra_data: dict | None = None,
required_capabilities: dict | None = None,
) -> RequestCandidate:
"""
创建候选记录
Args:
db: 数据库会话
request_id: 请求ID
candidate_index: 候选序号
retry_index: 重试序号从0开始
user_id: 用户ID
api_key_id: API Key ID
username: 用户名快照
api_key_name: API Key 名称快照
provider_id: Provider ID
endpoint_id: Endpoint ID
key_id: API Key ID
status: 候选状态 ('available', 'used', 'skipped', 'success', 'failed')
skip_reason: 跳过原因
is_cached: 是否为缓存亲和性候选
extra_data: 额外数据
required_capabilities: 请求需要的能力标签
"""
candidate = RequestCandidate(
id=str(candidate_id or uuid.uuid4()),
request_id=request_id,
candidate_index=candidate_index,
retry_index=retry_index, # 新增
user_id=user_id,
api_key_id=api_key_id,
username=username,
api_key_name=api_key_name,
provider_id=provider_id,
endpoint_id=endpoint_id,
key_id=key_id,
status=status,
skip_reason=skip_reason,
is_cached=is_cached,
extra_data=extra_data or {},
required_capabilities=required_capabilities,
created_at=datetime.now(timezone.utc),
)
db.add(candidate)
db.flush() # 只flush不立即 commit
# 标记为批量提交(非关键数据,可延迟)
get_batch_committer().mark_dirty(db)
return candidate
@staticmethod
def mark_candidate_started(db: Session, candidate_id: str) -> None:
"""
标记候选开始执行
Args:
db: 数据库会话
candidate_id: 候选ID
"""
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
if candidate:
candidate.status = "pending"
candidate.started_at = datetime.now(timezone.utc)
# 中间态改为 flush最终 success/failed 仍会立即提交,
# 但开始执行这一跳不再单独制造一次事务往返。
RequestCandidateService._persist_candidate_update(db, immediate=False)
@staticmethod
def update_candidate_status(db: Session, candidate_id: str, status: str) -> None:
"""
更新候选状态(通用方法)
Args:
db: 数据库会话
candidate_id: 候选ID
status: 新状态pending, available, success, failed, skipped
"""
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
if candidate:
candidate.status = status
# 如果状态变更为 pending记录开始时间
if status == "pending" and not candidate.started_at:
candidate.started_at = datetime.now(timezone.utc)
RequestCandidateService._persist_candidate_update(
db, immediate=status not in {"pending", "streaming"}
)
@staticmethod
def mark_candidate_streaming(
db: Session,
candidate_id: str,
concurrent_requests: int | None = None,
) -> None:
"""
标记候选为流式传输中
用于流式请求:连接建立成功后,流开始传输时调用。
此时请求尚未完成,需要等流传输完毕后再调用 mark_candidate_success。
注意streaming 阶段不设置 status_code最终状态码由
mark_candidate_success / mark_candidate_failed 在流结束时写入。
Args:
db: 数据库会话
candidate_id: 候选ID
concurrent_requests: 并发请求数
"""
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
if candidate:
candidate.status = "streaming"
candidate.concurrent_requests = concurrent_requests
# streaming 状态不设置 finished_at 和 status_code因为请求还在进行中
RequestCandidateService._persist_candidate_update(db, immediate=False)
@staticmethod
def mark_candidate_success(
db: Session,
candidate_id: str,
status_code: int,
latency_ms: int,
concurrent_requests: int | None = None,
extra_data: dict | None = None,
) -> None:
"""
标记候选执行成功
Args:
db: 数据库会话
candidate_id: 候选ID
status_code: HTTP 状态码
latency_ms: 延迟(毫秒)
concurrent_requests: 并发请求数
extra_data: 额外数据
"""
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
if candidate:
candidate.status = "success"
candidate.status_code = status_code
candidate.latency_ms = latency_ms
candidate.concurrent_requests = concurrent_requests
candidate.finished_at = datetime.now(timezone.utc)
# 成功时清空错误字段(可能是整流重试后成功,之前记录过错误)
candidate.error_type = None
candidate.error_message = None
if extra_data:
candidate.extra_data = {**(candidate.extra_data or {}), **extra_data}
# 关键状态更新:立即提交,不使用批量提交
# 原因:前端需要实时看到请求成功/失败状态
db.commit()
@staticmethod
def mark_candidate_failed(
db: Session,
candidate_id: str,
error_type: str,
error_message: str,
status_code: int | None = None,
latency_ms: int | None = None,
concurrent_requests: int | None = None,
extra_data: dict | None = None,
) -> None:
"""
标记候选执行失败
Args:
db: 数据库会话
candidate_id: 候选ID
error_type: 错误类型
error_message: 错误消息
status_code: HTTP 状态码(如果有)
latency_ms: 延迟(毫秒)
concurrent_requests: 并发请求数
extra_data: 额外数据
"""
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
if candidate:
candidate.status = "failed"
candidate.error_type = error_type
candidate.error_message = error_message
candidate.status_code = status_code
candidate.latency_ms = latency_ms
candidate.concurrent_requests = concurrent_requests
candidate.finished_at = datetime.now(timezone.utc)
if extra_data:
candidate.extra_data = {**(candidate.extra_data or {}), **extra_data}
# 关键状态更新:立即提交,不使用批量提交
# 原因:前端需要实时看到请求成功/失败状态
db.commit()
@staticmethod
def mark_candidate_cancelled(
db: Session,
candidate_id: str,
status_code: int = 499,
latency_ms: int | None = None,
concurrent_requests: int | None = None,
extra_data: dict | None = None,
) -> None:
"""
标记候选被客户端取消
客户端主动断开连接不算系统失败,使用 cancelled 状态。
Args:
db: 数据库会话
candidate_id: 候选ID
status_code: HTTP 状态码(通常是 499
latency_ms: 延迟(毫秒)
concurrent_requests: 并发请求数
extra_data: 额外数据
"""
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
if candidate:
candidate.status = "cancelled"
candidate.status_code = status_code
candidate.latency_ms = latency_ms
candidate.concurrent_requests = concurrent_requests
candidate.finished_at = datetime.now(timezone.utc)
if extra_data:
candidate.extra_data = {**(candidate.extra_data or {}), **extra_data}
db.commit()
@staticmethod
def mark_candidate_skipped(
db: Session,
candidate_id: str,
skip_reason: str | None = None,
*,
status_code: int | None = None,
concurrent_requests: int | None = None,
extra_data: dict | None = None,
) -> None:
"""
标记候选为已跳过
Args:
db: 数据库会话
candidate_id: 候选ID
skip_reason: 跳过原因
status_code: HTTP 状态码(可选)
concurrent_requests: 并发请求数(这里实际记录 RPM 计数)
extra_data: 额外数据(合并写入)
"""
candidate = db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
if candidate:
candidate.status = "skipped"
candidate.skip_reason = skip_reason
candidate.finished_at = datetime.now(timezone.utc)
if status_code is not None:
candidate.status_code = int(status_code)
if concurrent_requests is not None:
candidate.concurrent_requests = int(concurrent_requests)
if extra_data:
base = candidate.extra_data if isinstance(candidate.extra_data, dict) else {}
candidate.extra_data = {**base, **extra_data}
db.flush() # 只 flush不立即 commit
get_batch_committer().mark_dirty(db)
@staticmethod
def get_candidates_by_request_id(db: Session, request_id: str) -> list[RequestCandidate]:
"""
获取请求的所有候选记录
Args:
db: 数据库会话
request_id: 请求ID
Returns:
候选记录列表,按 candidate_index 排序
"""
return (
db.query(RequestCandidate)
.filter(RequestCandidate.request_id == request_id)
.order_by(RequestCandidate.candidate_index)
.all()
)
@staticmethod
def get_candidate_stats_by_provider(db: Session, provider_id: str, limit: int = 100) -> dict:
"""
获取 Provider 的候选统计
Args:
db: 数据库会话
provider_id: Provider ID
limit: 最近记录数量限制
Returns:
统计信息字典
"""
candidates = (
db.query(RequestCandidate)
.filter(RequestCandidate.provider_id == provider_id)
.order_by(RequestCandidate.created_at.desc())
.limit(limit)
.all()
)
total_candidates = len(candidates)
success_count = sum(1 for c in candidates if c.status == "success")
failed_count = sum(1 for c in candidates if c.status == "failed")
cancelled_count = sum(1 for c in candidates if c.status == "cancelled")
skipped_count = sum(1 for c in candidates if c.status == "skipped")
pending_count = sum(1 for c in candidates if c.status == "pending")
available_count = sum(1 for c in candidates if c.status == "available")
# 计算失败率只统计已完成的候选即成功或失败的cancelled 不算失败)
completed_count = success_count + failed_count
failure_rate = (failed_count / completed_count * 100) if completed_count > 0 else 0
return {
"total_attempts": total_candidates, # 前端使用 total_attempts 字段
"success_count": success_count,
"failed_count": failed_count,
"cancelled_count": cancelled_count, # 客户端取消数
"skipped_count": skipped_count,
"pending_count": pending_count,
"available_count": available_count, # 尚未被调度的候选数
"failure_rate": round(failure_rate, 2),
}
@staticmethod
def calculate_candidate_ttfb(
db: Session,
candidate_id: str,
request_start_time: float,
global_first_byte_time_ms: int,
) -> int:
"""
计算候选自身的首字节时间 (TTFB)
请求链路追踪中的 TTFB 应该是"该候选自身"的首字时间,
而不是整个请求从开始到收到首字节的时间。
Args:
db: 数据库会话
candidate_id: 候选 ID
request_start_time: 请求开始时间Unix timestamp
global_first_byte_time_ms: 全局首字节时间(相对于 request_start_time 的毫秒数)
Returns:
候选自身的 TTFB毫秒如果计算失败则返回 global_first_byte_time_ms
"""
try:
candidate = (
db.query(RequestCandidate).filter(RequestCandidate.id == candidate_id).first()
)
if candidate and candidate.started_at:
started_at = candidate.started_at
if started_at.tzinfo is None:
started_at = started_at.replace(tzinfo=timezone.utc)
# 使用整数毫秒计算,避免浮点精度问题
request_start_epoch_ms = round(request_start_time * 1000)
started_at_epoch_ms = round(started_at.timestamp() * 1000)
first_byte_epoch_ms = request_start_epoch_ms + global_first_byte_time_ms
return max(0, int(first_byte_epoch_ms - started_at_epoch_ms))
except Exception as e:
logger.debug("计算候选 TTFB 失败: {}", e)
return global_first_byte_time_ms