refactor: 拆分大型模块为 mixin/子模块结构

- cli_handler_base.py 拆分为 7 个 mixin (event/monitor/prefetch/request/sse_helpers/stream/sync)
- usage/service.py 拆分为 6 个子模块 (types/active_requests/cache_analysis/lifecycle/pricing/query/recording)
- models/database 拆分为独立模型文件 (auth/misc/model/provider/stats/usage/user)
- DUMMY_THOUGHT_SIGNATURE 常量提升到 core/api_format/conversion/constants 统一管理
- task/service.py 内联导入提升为顶层导入
- 流处理函数签名移除冗余的 http_client 参数
This commit is contained in:
fawney19
2026-02-14 12:55:58 +08:00
parent f464f32e48
commit 26ede849e2
32 changed files with 10173 additions and 7218 deletions

View File

@@ -0,0 +1,522 @@
"""CLI Handler - SSE 事件处理 + 格式转换 Mixin"""
from __future__ import annotations
import json
from typing import Any
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import get_format_converter_registry
from src.core.logger import logger
from src.services.provider.behavior import get_provider_behavior
from .cli_sse_helpers import (
_format_converted_events_to_sse,
_parse_gemini_json_array_line,
_parse_sse_data_line,
_parse_sse_event_data_line,
)
class CliEventMixin:
"""SSE 事件处理和格式转换相关方法的 Mixin"""
def _handle_sse_event(
self,
ctx: StreamContext,
event_name: str | None,
data_str: str,
record_chunk: bool = False,
) -> None:
"""
处理 SSE 事件
通用框架:解析 JSON、更新计数器
子类可覆盖 _process_event_data() 实现格式特定逻辑
Args:
ctx: 流上下文
event_name: 事件名称(如 message_start, content_block_delta 等)
data_str: 事件数据字符串JSON 格式)
record_chunk: 是否记录到 parsed_chunks不需要格式转换时应为 True
当为 True 时,同时更新 data_count
当为 False 时data_count 由 _record_converted_chunks 更新
"""
if not data_str:
return
if data_str == "[DONE]":
ctx.has_completion = True
return
try:
data = json.loads(data_str)
except json.JSONDecodeError:
return
if not isinstance(data, dict):
return
behavior = get_provider_behavior(
provider_type=ctx.provider_type,
endpoint_sig=ctx.provider_api_format,
)
envelope = behavior.envelope
if envelope:
data = envelope.unwrap_response(data)
if not isinstance(data, dict):
return
# 当不需要格式转换时,更新 data_count需要记录时再写入 parsed_chunks。
# 当需要格式转换时record_chunk=Falsedata_count 由 _record_converted_chunks 更新
if record_chunk:
ctx.data_count += 1
if ctx.record_parsed_chunks:
ctx.parsed_chunks.append(data)
event_type = event_name or data.get("type", "")
if envelope:
envelope.postprocess_unwrapped_response(model=ctx.model, data=data)
# 调用格式特定的处理逻辑
# 注意跨格式转换时_process_event_data 会自动选择正确的 Provider 解析器
self._process_event_data(ctx, event_type, data)
def _process_event_data(
self,
ctx: StreamContext,
event_type: str,
data: dict[str, Any],
) -> None:
"""
处理解析后的事件数据 - 子类应覆盖此方法
默认实现使用 ResponseParser 提取 usage
"""
# 提取 response_id
if not ctx.response_id:
response_obj = data.get("response")
if isinstance(response_obj, dict) and response_obj.get("id"):
ctx.response_id = response_obj["id"]
elif "id" in data:
ctx.response_id = data["id"]
# 使用解析器提取 usage
# Claude/CLI 流式响应的 usage 可能在首个 chunk 或最后一个 chunk 中
# 首个 chunk 可能部分为 0最后一个 chunk 包含完整值,因此取最大值确保正确计费
#
# 重要:当跨格式转换时,收到的数据是 Provider 格式,需要使用 Provider 格式的解析器
# 而不是客户端格式的解析器self.parser
parser = self.parser
if ctx.provider_api_format and ctx.provider_api_format != ctx.client_api_format:
# 跨格式转换:使用 Provider 格式的解析器
try:
provider_parser = get_parser_for_format(ctx.provider_api_format)
if provider_parser:
parser = provider_parser
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] 使用 Provider 解析器: "
f"{ctx.provider_api_format} (client={ctx.client_api_format})"
)
except KeyError:
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] 未找到 Provider 格式解析器: "
f"{ctx.provider_api_format}, 回退使用客户端格式解析器"
)
usage = parser.extract_usage_from_response(data)
if usage:
new_input = usage.get("input_tokens", 0)
new_output = usage.get("output_tokens", 0)
new_cached = usage.get("cache_read_tokens", 0)
new_cache_creation = usage.get("cache_creation_tokens", 0)
# 取最大值更新
if new_input > ctx.input_tokens:
ctx.input_tokens = new_input
if new_output > ctx.output_tokens:
ctx.output_tokens = new_output
if new_cached > ctx.cached_tokens:
ctx.cached_tokens = new_cached
if new_cache_creation > ctx.cache_creation_tokens:
ctx.cache_creation_tokens = new_cache_creation
# 保存最后一个非空 usage 作为 final_usage
if any([new_input, new_output, new_cached, new_cache_creation]):
ctx.final_usage = usage
# 提取文本内容(同样使用正确的解析器)
text = parser.extract_text_content(data)
if text:
ctx.append_text(text)
# 检查完成事件
if event_type in ("response.completed", "message_stop"):
ctx.has_completion = True
response_obj = data.get("response")
if isinstance(response_obj, dict):
ctx.final_response = response_obj
def _record_converted_chunks(
self,
ctx: StreamContext,
converted_events: list[dict[str, Any]],
) -> None:
"""
记录转换后的 chunk 数据到 parsed_chunks并更新统计信息
当需要格式转换时,记录的是转换后的数据(客户端实际收到的格式);
同时更新 data_count、has_completion 等统计信息。
重要:此方法也从转换后的事件中提取 usage 信息,作为 _process_event_data
从原始数据提取的补充。这确保即使原始 Provider 数据中没有 usage如 OpenAI
未设置 stream_options也能从转换后的格式中获取。
Args:
ctx: 流上下文
converted_events: 转换后的事件列表
"""
for evt in converted_events:
if isinstance(evt, dict):
ctx.data_count += 1
if ctx.record_parsed_chunks:
ctx.parsed_chunks.append(evt)
# 检测完成事件(根据客户端格式判断)
# OpenAI 格式: choices[].finish_reason
# Claude 格式: type == "message_stop" 或 stop_reason
event_type = evt.get("type", "")
if event_type == "message_stop":
ctx.has_completion = True
elif event_type == "response.completed":
ctx.has_completion = True
elif "choices" in evt:
choices = evt.get("choices", [])
for choice in choices:
if isinstance(choice, dict) and choice.get("finish_reason"):
ctx.has_completion = True
break
# 从转换后的事件中提取 usage补充 _process_event_data 的提取)
# Claude 格式: message_delta.usage 或 message_start.message.usage
# OpenAI 格式: chunk.usage
self._extract_usage_from_converted_event(ctx, evt, event_type)
def _extract_usage_from_converted_event(
self,
ctx: StreamContext,
evt: dict[str, Any],
event_type: str,
) -> None:
"""
从转换后的事件中提取 usage 信息
支持多种格式:
- Claude: message_delta.usage, message_start.message.usage
- OpenAI: chunk.usage
- Gemini: usageMetadata
Args:
ctx: 流上下文
evt: 转换后的事件
event_type: 事件类型
"""
usage: dict[str, Any] | None = None
# Claude 格式: message_delta 或 message_start
if event_type == "message_delta":
usage = evt.get("usage")
elif event_type == "message_start":
message = evt.get("message", {})
if isinstance(message, dict):
usage = message.get("usage")
# OpenAI Responses API (openai:cli) 格式: response.completed 中 usage 嵌套在 response 对象内
elif event_type == "response.completed":
resp_obj = evt.get("response")
if isinstance(resp_obj, dict):
usage = resp_obj.get("usage")
# 兼容: 部分实现可能在顶层也有 usage
if not usage:
usage = evt.get("usage")
# OpenAI Chat 格式: 直接在 chunk 中
elif "usage" in evt:
usage = evt.get("usage")
# Gemini 格式: usageMetadata
elif "usageMetadata" in evt:
meta = evt.get("usageMetadata", {})
if isinstance(meta, dict):
usage = {
"input_tokens": meta.get("promptTokenCount", 0),
"output_tokens": meta.get("candidatesTokenCount", 0),
"cache_read_tokens": meta.get("cachedContentTokenCount", 0),
"cache_creation_tokens": 0, # Gemini 目前不支持缓存创建
}
if usage and isinstance(usage, dict):
new_input = usage.get("input_tokens", 0) or 0
new_output = usage.get("output_tokens", 0) or 0
new_cached = usage.get("cache_read_tokens") or usage.get("cache_read_input_tokens") or 0
new_cache_creation = (
usage.get("cache_creation_tokens") or usage.get("cache_creation_input_tokens") or 0
)
# 取最大值更新(与 _process_event_data 相同的策略)
if new_input > ctx.input_tokens:
ctx.input_tokens = new_input
logger.debug("[{}] 从转换后事件更新 input_tokens: {}", ctx.request_id, new_input)
if new_output > ctx.output_tokens:
ctx.output_tokens = new_output
logger.debug("[{}] 从转换后事件更新 output_tokens: {}", ctx.request_id, new_output)
if new_cached > ctx.cached_tokens:
ctx.cached_tokens = new_cached
if new_cache_creation > ctx.cache_creation_tokens:
ctx.cache_creation_tokens = new_cache_creation
# 保存最后一个非空 usage
if any([new_input, new_output, new_cached, new_cache_creation]):
ctx.final_usage = usage
def _finalize_stream_metadata(self, ctx: StreamContext) -> None:
"""
在记录统计前从 parsed_chunks 中提取额外的元数据 - 子类可覆盖
这是一个后处理钩子,在流传输完成后、记录 Usage 之前调用。
子类可以覆盖此方法从 ctx.parsed_chunks 中提取格式特定的元数据,
如 Gemini 的 modelVersion、token 统计等。
Args:
ctx: 流上下文,包含 parsed_chunks 和 response_metadata
"""
pass
def _needs_format_conversion(self, ctx: StreamContext) -> bool:
"""
[已废弃] 仅根据格式差异判断是否需要转换
警告:此方法只检查格式是否不同,不检查端点的 format_acceptance_config 配置!
正确的判断应使用候选筛选阶段的结果ctx.needs_conversion该结果由
is_format_compatible() 函数根据全局开关和端点配置计算得出。
此方法保留仅供调试和日志输出使用,流生成器中不应调用此方法。
当 Provider 的 API 格式与客户端请求的 API 格式不同时,需要转换响应。
例如:客户端请求 Claude 格式,但 Provider 返回 OpenAI 格式。
注意:
- CLAUDE 和 CLAUDE_CLI、GEMINI 和 GEMINI_CLI格式相同只是认证不同可透传
- OPENAI 和 OPENAI_CLI格式不同Chat Completions vs Responses API需要转换
"""
from src.core.api_format.metadata import can_passthrough_endpoint
from src.core.api_format.signature import normalize_signature_key
if not ctx.provider_api_format or not ctx.client_api_format:
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
f"provider_api_format={ctx.provider_api_format!r}, client_api_format={ctx.client_api_format!r} -> False (missing)"
)
return False
provider_format = normalize_signature_key(str(ctx.provider_api_format))
client_format = normalize_signature_key(str(ctx.client_api_format))
# 1. 格式完全匹配 -> 不需要转换
if provider_format == client_format:
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
f"provider={provider_format}, client={client_format} -> False (exact match)"
)
return False
# 2. 根据 data_format_id 判断是否可透传(可透传则不需要转换)
if can_passthrough_endpoint(client_format, provider_format):
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
f"provider={provider_format}, client={client_format} -> False (passthroughable)"
)
return False
# 3. 其他情况 -> 需要转换
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
f"provider={provider_format}, client={client_format} -> True"
)
return True
def _mark_first_output(self, ctx: StreamContext, state: dict[str, bool]) -> None:
"""
标记首次输出:记录 TTFB 并更新 streaming 状态
在第一次 yield 数据前调用,确保:
1. 首字时间 (TTFB) 已记录到 ctx
2. Usage 状态已更新为 streaming包含 provider/key/TTFB 信息)
Args:
ctx: 流上下文
state: 包含 first_yield 和 streaming_updated 的状态字典
"""
if state["first_yield"]:
ctx.record_first_byte_time(self.start_time)
state["first_yield"] = False
if not state["streaming_updated"]:
# 优先使用当前请求的 DB 会话同步更新,避免状态延迟或丢失
try:
from src.services.usage import UsageService
UsageService.update_usage_status(
db=self.db,
request_id=self.request_id,
status="streaming",
provider=ctx.provider_name,
target_model=ctx.mapped_model,
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
first_byte_time_ms=ctx.first_byte_time_ms,
api_format=ctx.api_format,
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
)
except Exception as e:
logger.warning("[{}] 同步更新 streaming 状态失败: {}", self.request_id, e)
# 回退到后台任务更新
self._update_usage_to_streaming_with_ctx(ctx)
state["streaming_updated"] = True
def _convert_sse_line(
self,
ctx: StreamContext,
line: str,
events: list, # noqa: ARG002 - 预留给上下文感知转换
) -> tuple[list[str], list[dict[str, Any]]]:
"""
将 SSE 行从 Provider 格式转换为客户端格式
Args:
ctx: 流上下文
line: 原始 SSE 行
events: 当前累积的事件列表(预留参数,用于未来上下文感知转换如合并相邻事件)
Returns:
(sse_lines, converted_events) 元组:
- sse_lines: 转换后的 SSE 行列表(一入多出),空列表表示跳过该行
- converted_events: 转换后的事件对象列表(用于记录到 parsed_chunks
"""
# 空行直接返回
if not line or line.strip() == "":
return ([line] if line else [], [])
client_format = (ctx.client_api_format or "").strip().lower()
# [DONE] 标记处理:只有 OpenAI 客户端需要Claude 客户端不需要
if line == "data: [DONE]":
if client_format.startswith("openai"):
return [line], []
else:
# Claude/Gemini 客户端不需要 [DONE] 标记
return [], []
provider_format = (ctx.provider_api_format or "").strip().lower()
# 过滤上游控制行id/retry避免与目标格式混淆
if line.startswith(("id:", "retry:")):
return [], []
# 解析 SSE 行为 JSON 对象
data_obj, status = self._parse_sse_line_to_json(line, provider_format)
# 根据解析状态决定行为
if status == "empty" or status == "skip":
return [], []
if status == "invalid" or status == "passthrough":
return [line], []
behavior = get_provider_behavior(
provider_type=ctx.provider_type,
endpoint_sig=ctx.provider_api_format,
)
envelope = behavior.envelope
if envelope:
data_obj = envelope.unwrap_response(data_obj)
envelope.postprocess_unwrapped_response(model=ctx.model, data=data_obj)
# 初始化流式转换状态
if ctx.stream_conversion_state is None:
from src.core.api_format.conversion.stream_state import StreamState
# 使用客户端请求的模型ctx.model而非映射后的上游模型ctx.mapped_model
init_model = ctx.model or ""
logger.debug(
f"[{ctx.request_id}] StreamState init: ctx.model={ctx.model!r}, "
f"mapped_model={ctx.mapped_model!r}, using={init_model!r}"
)
ctx.stream_conversion_state = StreamState(
model=init_model,
message_id=ctx.response_id or ctx.request_id or "",
)
# 执行格式转换
try:
registry = get_format_converter_registry()
# status == "ok" 时 data_obj 必定是有效的 dict防御性检查
if data_obj is None:
return [], []
converted_events = registry.convert_stream_chunk(
data_obj,
provider_format,
client_format,
state=ctx.stream_conversion_state,
)
result = _format_converted_events_to_sse(converted_events, client_format)
if result:
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] 流式转换: "
f"{provider_format}->{client_format}, events={len(converted_events)}, "
f"first_output={result[0][:100] if result else 'empty'}..."
)
return result, converted_events
except Exception as e:
logger.warning("格式转换失败,透传原始数据: {}", e)
return [line], []
def _parse_sse_line_to_json(self, line: str, provider_format: str) -> tuple[Any | None, str]:
"""
解析 SSE 行为 JSON 对象
支持多种格式:
- 标准 SSE: "data: {...}"
- event+data 同行: "event: xxx data: {...}"
- Gemini JSON-array: 裸 JSON 行
Args:
line: 原始 SSE 行
provider_format: Provider API 格式
Returns:
(parsed_json, status) 元组:
- (obj, "ok") - 解析成功
- (None, "empty") - 内容为空,应跳过
- (None, "invalid") - JSON 解析失败,应透传原始行
- (None, "skip") - 应跳过(如纯 event 行)
- (None, "passthrough") - 无法识别,应透传原始行
"""
# 标准 SSE: data: {...}
if line.startswith("data:"):
return _parse_sse_data_line(line)
# event + data 同行: event: xxx data: {...}
if line.startswith("event:") and " data:" in line:
return _parse_sse_event_data_line(line)
# 纯 event 行不参与转换
if line.startswith("event:"):
return None, "skip"
# Gemini JSON-array 格式
if provider_format.startswith("gemini"):
return _parse_gemini_json_array_line(line)
# 其他格式:无法识别,透传
return None, "passthrough"

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,503 @@
"""CLI Handler - 监控/统计 Mixin"""
from __future__ import annotations
import asyncio
import time
from collections.abc import AsyncGenerator
from typing import Any
import httpx
from fastapi import Request
from src.api.handlers.base.base_handler import MessageTelemetry
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import filter_proxy_response_headers
from src.core.error_utils import extract_client_error_message
from src.core.exceptions import (
ProviderAuthException,
ProviderRateLimitException,
ProviderTimeoutException,
ThinkingSignatureException,
)
from src.core.logger import logger
from src.database import get_db
from src.models.database import User
from src.services.provider.behavior import get_provider_behavior
class CliMonitorMixin:
"""监控和统计相关方法的 Mixin"""
async def _create_monitored_stream(
self,
ctx: StreamContext,
stream_generator: AsyncGenerator[bytes],
http_request: Request | None = None,
) -> AsyncGenerator[bytes]:
"""
创建带监控的流生成器
支持两种断连检测方式:
1. 如果提供了 http_request使用后台任务主动检测客户端断连
2. 如果未提供,仅依赖 asyncio.CancelledError 被动检测
Args:
ctx: 流上下文
stream_generator: 底层流生成器
http_request: FastAPI Request 对象,用于检测客户端断连
"""
import time as time_module
last_chunk_time = time_module.time()
chunk_count = 0
try:
if http_request is not None:
# 使用后台任务检测断连,完全不阻塞流式传输
disconnected = False
async def check_disconnect_background() -> None:
nonlocal disconnected
while not disconnected and not ctx.has_completion:
await asyncio.sleep(0.5)
try:
if await http_request.is_disconnected():
disconnected = True
break
except Exception as e:
# 检测失败时不中断流,继续传输
logger.debug("ID:{} | 断连检测异常: {}", ctx.request_id, e)
# 启动后台检查任务
check_task = asyncio.create_task(check_disconnect_background())
try:
async for chunk in stream_generator:
if disconnected:
# 如果响应已完成,客户端断开不算失败
if ctx.has_completion:
logger.info(
f"ID:{ctx.request_id} | Client disconnected after completion"
)
else:
logger.warning("ID:{} | Client disconnected", ctx.request_id)
ctx.status_code = 499
ctx.error_message = "client_disconnected"
break
last_chunk_time = time_module.time()
chunk_count += 1
yield chunk
finally:
check_task.cancel()
try:
await check_task
except asyncio.CancelledError:
pass
else:
# 无 http_request仅被动监控
async for chunk in stream_generator:
last_chunk_time = time_module.time()
chunk_count += 1
yield chunk
except asyncio.CancelledError:
# 注意CancelledError 不等于"用户手动取消",它既可能是客户端断连触发,
# 也可能是服务端(重载/关停/内部取消)导致的协程取消。
# 这里尽量做一次"断连归因":仅当能确认客户端已断开时才记为 499 cancelled。
time_since_last_chunk = time_module.time() - last_chunk_time
is_client_disconnected = False
if http_request is not None:
try:
# shield + timeout: 避免在取消态下二次被 CancelledError 打断,尽力取到断连状态
# 限时 0.5s 防止极端情况下的阻塞
is_client_disconnected = await asyncio.wait_for(
asyncio.shield(http_request.is_disconnected()),
timeout=0.5,
)
except (asyncio.CancelledError, asyncio.TimeoutError):
# 无法在取消态/超时下完成断连检查,保守视为未知(不强行归因为客户端)
is_client_disconnected = False
except Exception as e:
logger.debug("ID:{} | cancel 断连检测失败: {}", ctx.request_id, e)
is_client_disconnected = False
# 如果响应已完成,不标记为失败/取消
if not ctx.has_completion:
if is_client_disconnected:
ctx.status_code = 499
ctx.error_message = "client_disconnected"
logger.warning(
f"ID:{ctx.request_id} | Stream cancelled by client: "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}"
)
else:
# 服务端中断(例如重载/关停/内部取消) -- 不应伪装成客户端取消
ctx.status_code = 503
ctx.error_message = "server_cancelled"
logger.error(
f"ID:{ctx.request_id} | Stream interrupted by server: "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}"
)
raise
except httpx.TimeoutException as e:
ctx.status_code = 504
ctx.error_message = str(e)
raise
except Exception as e:
ctx.status_code = 500
ctx.error_message = str(e)
raise
async def _record_stream_stats(
self,
ctx: StreamContext,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""在流完成后记录统计信息"""
try:
# 使用 self.start_time 作为时间基准,与首字时间保持一致
# 注意:不要把统计延迟算进响应时间里
response_time_ms = int((time.time() - self.start_time) * 1000)
await asyncio.sleep(0.1)
if not ctx.provider_name:
logger.warning("[{}] 流式请求失败,未选中提供商", ctx.request_id)
return
behavior = get_provider_behavior(
provider_type=ctx.provider_type,
endpoint_sig=ctx.provider_api_format,
)
envelope = behavior.envelope
if envelope:
envelope.on_http_status(
base_url=ctx.selected_base_url,
status_code=ctx.status_code,
)
# 获取新的 DB session
db_gen = get_db()
bg_db = next(db_gen)
try:
from src.models.database import ApiKey as ApiKeyModel
user = bg_db.query(User).filter(User.id == ctx.user_id).first()
api_key = bg_db.query(ApiKeyModel).filter(ApiKeyModel.id == ctx.api_key_id).first()
if not user or not api_key:
logger.warning(
"[{}] 无法记录统计: user={} api_key={}",
ctx.request_id,
user is not None,
api_key is not None,
)
return
bg_telemetry = MessageTelemetry(
bg_db, user, api_key, ctx.request_id, self.client_ip
)
response_body = {
"chunks": ctx.parsed_chunks,
"metadata": {
"stream": True,
"total_chunks": len(ctx.parsed_chunks),
"data_count": ctx.data_count,
"has_completion": ctx.has_completion,
"response_time_ms": response_time_ms,
},
}
# 使用实际发送给 Provider 的请求体(如果有),否则用原始请求体
actual_request_body = ctx.provider_request_body or original_request_body
# 根据状态码决定记录成功还是失败
# 499 = 客户端取消(不算系统失败);其他 4xx/5xx 视为失败
if ctx.status_code and ctx.status_code >= 400:
client_response_headers = ctx.client_response_headers or {
"content-type": "application/json"
}
if ctx.is_client_disconnected():
# 客户端取消:记录为 cancelled不算系统失败
request_metadata = {"perf": ctx.perf_metrics} if ctx.perf_metrics else None
await bg_telemetry.record_cancelled(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
first_byte_time_ms=ctx.first_byte_time_ms,
status_code=ctx.status_code,
request_headers=original_headers,
request_body=actual_request_body,
is_stream=True,
api_format=ctx.api_format,
provider_request_headers=ctx.provider_request_headers,
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
response_body=response_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
target_model=ctx.mapped_model,
request_metadata=request_metadata,
)
logger.debug("{} 流式响应被客户端取消", self.FORMAT_ID)
logger.info(
f"[CANCEL] {self.request_id[:8]} | {ctx.model} | {ctx.provider_name} | {response_time_ms}ms | "
f"{ctx.status_code} | in:{ctx.input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}"
)
else:
# 服务端/上游异常:记录为失败
request_metadata = {"perf": ctx.perf_metrics} if ctx.perf_metrics else None
await bg_telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
status_code=ctx.status_code,
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
request_headers=original_headers,
request_body=actual_request_body,
is_stream=True,
api_format=ctx.api_format,
provider_request_headers=ctx.provider_request_headers,
# 预估 token 信息(来自 message_start 事件)
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
response_body=response_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
# 模型映射信息
target_model=ctx.mapped_model,
request_metadata=request_metadata,
)
logger.debug("{} 流式响应中断", self.FORMAT_ID)
logger.info(
f"[FAIL] {self.request_id[:8]} | {ctx.model} | {ctx.provider_name} | {response_time_ms}ms | "
f"{ctx.status_code} | in:{ctx.input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}"
)
else:
# 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据
self._finalize_stream_metadata(ctx)
# 流未正常完成(如上游截断/连接中断)且无 token 数据时,
# 从已收集的文本和请求体估算 tokens避免 usage 记录为 0
if (
not ctx.has_completion
and ctx.data_count > 0
and ctx.input_tokens == 0
and ctx.output_tokens == 0
):
self._estimate_tokens_for_incomplete_stream(ctx, actual_request_body)
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
client_response_headers.update(
{
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
"content-type": "text/event-stream",
}
)
logger.debug(
f"[{ctx.request_id}] 开始记录 Usage: "
f"provider={ctx.provider_name}, model={ctx.model}, "
f"in={ctx.input_tokens}, out={ctx.output_tokens}"
)
request_metadata = {"perf": ctx.perf_metrics} if ctx.perf_metrics else None
total_cost = await bg_telemetry.record_success(
provider=ctx.provider_name,
model=ctx.model,
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
response_time_ms=response_time_ms,
first_byte_time_ms=ctx.first_byte_time_ms, # 传递首字时间
status_code=ctx.status_code,
request_headers=original_headers,
request_body=actual_request_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
response_body=response_body,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
is_stream=True,
provider_request_headers=ctx.provider_request_headers,
api_format=ctx.api_format,
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
# Provider 侧追踪信息(用于记录真实成本)
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
# 模型映射信息
target_model=ctx.mapped_model,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata=ctx.response_metadata if ctx.response_metadata else None,
request_metadata=request_metadata,
)
logger.debug("[{}] Usage 记录完成: cost=${:.6f}", ctx.request_id, total_cost)
# 简洁的请求完成摘要(两行格式)
ttfb_part = (
f" | TTFB: {ctx.first_byte_time_ms}ms" if ctx.first_byte_time_ms else ""
)
logger.info(
"[OK] {} | {} | {}{}\n Total: {}ms | in:{} out:{}",
self.request_id[:8],
ctx.model,
ctx.provider_name,
ttfb_part,
response_time_ms,
ctx.input_tokens or 0,
ctx.output_tokens or 0,
)
# 更新候选记录的最终状态和延迟时间
# 注意RequestExecutor 会在流开始时过早地标记成功(只记录了连接建立的时间)
# 这里用流传输完成后的实际时间覆盖
if ctx.attempt_id:
from src.services.request.candidate import RequestCandidateService
# 计算候选自身的 TTFB
candidate_first_byte_time_ms: int | None = None
if ctx.first_byte_time_ms is not None:
candidate_first_byte_time_ms = (
RequestCandidateService.calculate_candidate_ttfb(
db=bg_db,
candidate_id=ctx.attempt_id,
request_start_time=self.start_time,
global_first_byte_time_ms=ctx.first_byte_time_ms,
)
)
# 根据状态码决定是成功还是失败
# 499 = 客户端断开连接,应标记为失败
# 503 = 服务不可用(如流中断),应标记为失败
if ctx.status_code and ctx.status_code >= 400:
# 请求链路追踪使用 upstream_response原始响应回退到 error_message友好消息
trace_error_message = (
ctx.upstream_response or ctx.error_message or f"HTTP {ctx.status_code}"
)
extra_data = {
"stream_completed": False,
"chunk_count": ctx.chunk_count,
"data_count": ctx.data_count,
}
if ctx.proxy_info:
extra_data["proxy"] = ctx.proxy_info
if candidate_first_byte_time_ms is not None:
extra_data["first_byte_time_ms"] = candidate_first_byte_time_ms
if ctx.is_client_disconnected():
RequestCandidateService.mark_candidate_cancelled(
db=bg_db,
candidate_id=ctx.attempt_id,
status_code=ctx.status_code,
latency_ms=response_time_ms,
extra_data=extra_data,
)
else:
RequestCandidateService.mark_candidate_failed(
db=bg_db,
candidate_id=ctx.attempt_id,
error_type="stream_error",
error_message=trace_error_message,
status_code=ctx.status_code,
latency_ms=response_time_ms,
extra_data=extra_data,
)
else:
extra_data = {
"stream_completed": True,
"chunk_count": ctx.chunk_count,
"data_count": ctx.data_count,
}
if ctx.proxy_info:
extra_data["proxy"] = ctx.proxy_info
if ctx.rectified:
extra_data["rectified"] = True
if candidate_first_byte_time_ms is not None:
extra_data["first_byte_time_ms"] = candidate_first_byte_time_ms
RequestCandidateService.mark_candidate_success(
db=bg_db,
candidate_id=ctx.attempt_id,
status_code=ctx.status_code,
latency_ms=response_time_ms,
extra_data=extra_data,
)
finally:
bg_db.close()
except Exception as e:
logger.exception("记录流式统计信息时出错")
async def _record_stream_failure(
self,
ctx: StreamContext,
error: Exception,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""记录流式请求失败"""
# 使用 self.start_time 作为时间基准,与首字时间保持一致
response_time_ms = int((time.time() - self.start_time) * 1000)
status_code = 503
if isinstance(error, ThinkingSignatureException):
status_code = 400
elif isinstance(error, ProviderAuthException):
status_code = 503
elif isinstance(error, ProviderRateLimitException):
status_code = 429
elif isinstance(error, ProviderTimeoutException):
status_code = 504
ctx.status_code = status_code
ctx.error_message = str(error)
# 使用实际发送给 Provider 的请求体(如果有),否则用原始请求体
actual_request_body = ctx.provider_request_body or original_request_body
# 失败时返回给客户端的是 JSON 错误响应
client_response_headers = {"content-type": "application/json"}
request_metadata = {"perf": ctx.perf_metrics} if ctx.perf_metrics else None
await self.telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
status_code=status_code,
error_message=extract_client_error_message(error),
request_headers=original_headers,
request_body=actual_request_body,
is_stream=True,
api_format=ctx.api_format,
provider_request_headers=ctx.provider_request_headers,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
# 模型映射信息
target_model=ctx.mapped_model,
request_metadata=request_metadata,
)

View File

@@ -0,0 +1,610 @@
"""CLI Handler - Prefetch 和错误检测 Mixin"""
from __future__ import annotations
import codecs
import json
import time
from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING, Any
import httpx
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import (
check_html_response,
check_prefetched_response_error,
)
from src.config.constants import StreamDefaults
from src.config.settings import config
from src.core.exceptions import (
EmbeddedErrorException,
ProviderNotAvailableException,
ProviderTimeoutException,
)
from src.core.logger import logger
from src.services.provider.behavior import get_provider_behavior
from src.utils.sse_parser import SSEEventParser
from src.utils.timeout import read_first_chunk_with_ttfb_timeout
if TYPE_CHECKING:
from src.models.database import Provider, ProviderEndpoint
class CliPrefetchMixin:
"""Prefetch 和错误检测相关方法的 Mixin"""
def _flush_remaining_sse_data(
self,
ctx: StreamContext,
buffer: bytes,
decoder: codecs.IncrementalDecoder,
sse_parser: SSEEventParser,
*,
record_chunk: bool = True,
) -> None:
"""
异常发生时 flush 残留的字节 buffer 和 SSE parser 内部缓冲区。
用于 StreamClosed / RemoteProtocolError 等场景:
连接断开可能恰好发生在最后一个 SSE 事件(如 response.completed
的 data 行已收到、但终止空行尚未到达之时。此方法确保这些事件仍能被处理,
从而正确捕获 usage 等关键信息。
"""
try:
# 1) flush 字节 buffer 中的残余行
if buffer:
remaining = decoder.decode(buffer, True)
for line in remaining.split("\n"):
stripped = line.rstrip("\r")
events = sse_parser.feed_line(stripped)
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=record_chunk,
)
# 2) flush SSE parser 内部累积的未完成事件
for event in sse_parser.flush():
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=record_chunk,
)
except Exception:
# best-effort: 不应因 flush 失败影响后续流程
pass
def _estimate_tokens_for_incomplete_stream(
self,
ctx: StreamContext,
request_body: dict[str, Any],
) -> None:
"""
流未正常完成(无 response.completed且 token 均为 0 时的兜底估算。
从已收集的输出文本和请求体粗略估算 token 数,确保 usage 记录不为 0。
估算采用 ~4 字符/token 的保守比例。
"""
# 输出 tokens从已收集的文本估算
collected = ctx.collected_text
if collected:
ctx.output_tokens = max(1, len(collected) // 4)
# 输入 tokens从请求体文本内容估算
try:
total_input_len = 0
instructions = request_body.get("instructions")
if isinstance(instructions, str):
total_input_len += len(instructions)
# OpenAI Responses API 使用 input 字段Claude 使用 messages
input_items = request_body.get("input") or request_body.get("messages") or []
if isinstance(input_items, list):
for item in input_items:
if isinstance(item, str):
total_input_len += len(item)
elif isinstance(item, dict):
content = item.get("content", "")
if isinstance(content, str):
total_input_len += len(content)
elif isinstance(content, list):
for block in content:
if isinstance(block, dict):
text = block.get("text", "")
if isinstance(text, str):
total_input_len += len(text)
if total_input_len > 0:
ctx.input_tokens = max(1, total_input_len // 4)
else:
# fallback: 整个请求体 JSON 大小
body_str = json.dumps(request_body, ensure_ascii=False)
ctx.input_tokens = max(1, len(body_str) // 4)
except Exception:
pass
if ctx.input_tokens > 0 or ctx.output_tokens > 0:
logger.warning(
"[{}] 流未正常完成 (has_completion=False, data_count={}), "
"使用估算 tokens: in={}, out={}",
ctx.request_id,
ctx.data_count,
ctx.input_tokens,
ctx.output_tokens,
)
async def _prefetch_and_check_embedded_error(
self,
byte_iterator: Any,
provider: "Provider",
endpoint: "ProviderEndpoint",
ctx: StreamContext,
) -> list:
"""
预读流的前几行,检测嵌套错误
某些 Provider如 Gemini可能返回 HTTP 200但在响应体中包含错误信息。
这种情况需要在流开始输出之前检测,以便触发重试逻辑。
同时检测 HTML 响应(通常是 base_url 配置错误导致返回网页)。
首次读取时会应用 TTFB首字节超时检测超时则触发故障转移。
Args:
byte_iterator: 字节流迭代器
provider: Provider 对象
endpoint: Endpoint 对象
ctx: 流上下文
Returns:
预读的字节块列表(需要在后续流中先输出)
Raises:
EmbeddedErrorException: 如果检测到嵌套错误
ProviderNotAvailableException: 如果检测到 HTML 响应(配置错误)
ProviderTimeoutException: 如果首字节超时TTFB timeout
"""
prefetched_chunks: list = []
max_prefetch_lines = config.stream_prefetch_lines # 最多预读行数来检测错误
max_prefetch_bytes = StreamDefaults.MAX_PREFETCH_BYTES # 避免无换行响应导致 buffer 增长
total_prefetched_bytes = 0
buffer = b""
line_count = 0
should_stop = False
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
try:
# 获取对应格式的解析器
provider_format = ctx.provider_api_format
if provider_format:
try:
provider_parser = get_parser_for_format(provider_format)
except KeyError:
provider_parser = self.parser
else:
provider_parser = self.parser
# 使用共享的 TTFB 超时函数读取首字节
# 优先使用 Provider 配置,否则使用全局配置
ttfb_timeout = provider.stream_first_byte_timeout or config.stream_first_byte_timeout
first_chunk, aiter = await read_first_chunk_with_ttfb_timeout(
byte_iterator,
timeout=ttfb_timeout,
request_id=self.request_id,
provider_name=str(provider.name),
)
prefetched_chunks.append(first_chunk)
total_prefetched_bytes += len(first_chunk)
buffer += first_chunk
# 继续读取剩余的预读数据
async for chunk in aiter:
prefetched_chunks.append(chunk)
total_prefetched_bytes += len(chunk)
buffer += chunk
# 尝试按行解析缓冲区SSE 格式)
while b"\n" in buffer:
line_bytes, buffer = buffer.split(b"\n", 1)
try:
# 使用增量解码器,可以正确处理跨 chunk 的多字节字符
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
except Exception as e:
logger.warning(
f"[{self.request_id}] 预读时 UTF-8 解码失败: {e}, "
f"bytes={line_bytes[:50]!r}"
)
continue
line_count += 1
normalized_line = line.rstrip("\r")
# 检测 HTML 响应base_url 配置错误的常见症状)
if check_html_response(normalized_line):
logger.error(
f" [{self.request_id}] 检测到 HTML 响应,可能是 base_url 配置错误: "
f"Provider={provider.name}, Endpoint={endpoint.id[:8]}..., "
f"base_url={endpoint.base_url}"
)
raise ProviderNotAvailableException(
"上游服务返回了非预期的响应格式",
provider_name=str(provider.name),
upstream_status=200,
upstream_response=(
normalized_line[:500] if normalized_line else "(empty)"
),
)
if not normalized_line or normalized_line.startswith(":"):
# 空行或注释行,继续预读
if line_count >= max_prefetch_lines:
break
continue
# 尝试解析 SSE 数据
data_str = normalized_line
if normalized_line.startswith("data: "):
data_str = normalized_line[6:]
if data_str == "[DONE]":
should_stop = True
break
try:
data = json.loads(data_str)
except json.JSONDecodeError:
# 不是有效 JSON可能是部分数据继续
if line_count >= max_prefetch_lines:
break
continue
# 使用解析器检查是否为错误响应
if isinstance(data, dict) and provider_parser.is_error_response(data):
# 提取错误信息
parsed = provider_parser.parse_response(data, 200)
logger.warning(
f" [{self.request_id}] 检测到嵌套错误: "
f"Provider={provider.name}, "
f"error_type={parsed.error_type}, "
f"message={parsed.error_message}"
)
raise EmbeddedErrorException(
provider_name=str(provider.name),
error_code=(
int(parsed.error_type)
if parsed.error_type and parsed.error_type.isdigit()
else None
),
error_message=parsed.error_message,
error_status=parsed.error_type,
)
# 预读到有效数据,没有错误,停止预读
should_stop = True
break
# 达到预读字节上限,停止继续预读(避免无换行响应导致内存增长)
if not should_stop and total_prefetched_bytes >= max_prefetch_bytes:
logger.debug(
f" [{self.request_id}] 预读达到字节上限,停止继续预读: "
f"Provider={provider.name}, bytes={total_prefetched_bytes}, "
f"max_bytes={max_prefetch_bytes}"
)
break
if should_stop or line_count >= max_prefetch_lines:
break
# 预读结束后,检查是否为非 SSE 格式的 HTML/JSON 响应
# 处理某些代理返回的纯 JSON 错误(可能无换行/多行 JSON以及 HTML 页面base_url 配置错误)
if not should_stop and prefetched_chunks:
check_prefetched_response_error(
prefetched_chunks=prefetched_chunks,
parser=provider_parser,
request_id=self.request_id,
provider_name=str(provider.name),
endpoint_id=endpoint.id,
base_url=endpoint.base_url,
)
except (EmbeddedErrorException, ProviderTimeoutException, ProviderNotAvailableException):
# 重新抛出可重试的 Provider 异常,触发故障转移
raise
except OSError as e:
# 网络 I/O 异常:记录警告,可能需要重试
logger.warning(
" [{}] 预读流时发生网络异常: {}: {}", self.request_id, type(e).__name__, e
)
except Exception as e:
# 未预期的严重异常:记录错误并重新抛出,避免掩盖问题
logger.error(
f" [{self.request_id}] 预读流时发生严重异常: {type(e).__name__}: {e}",
exc_info=True,
)
raise
return prefetched_chunks
async def _create_response_stream_with_prefetch(
self,
ctx: StreamContext,
byte_iterator: Any,
response_ctx: Any,
prefetched_chunks: list,
) -> AsyncGenerator[bytes]:
"""创建响应流生成器(带预读数据,使用字节流)"""
try:
sse_parser = SSEEventParser()
last_data_time = time.time()
buffer = b""
output_state = {"first_yield": True, "streaming_updated": False}
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
# 使用已设置的 ctx.needs_conversion由候选筛选阶段根据端点配置判断
# 不再调用 _needs_format_conversion它只检查格式差异不检查端点配置
needs_conversion = ctx.needs_conversion
behavior = get_provider_behavior(
provider_type=ctx.provider_type,
endpoint_sig=ctx.provider_api_format,
)
envelope = behavior.envelope
if envelope and envelope.force_stream_rewrite():
needs_conversion = True
ctx.needs_conversion = True
# Kiro 特殊处理AWS Event Stream 二进制流需要重写为 SSE
ctx_provider_type = str(ctx.provider_type or "").strip().lower()
if ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
from src.services.provider.adapters.kiro.eventstream_rewriter import (
apply_kiro_stream_rewrite,
)
byte_iterator = apply_kiro_stream_rewrite(
byte_iterator,
model=str(ctx.model or ""),
input_tokens=int(ctx.input_tokens or 0),
prefetched_chunks=list(prefetched_chunks) if prefetched_chunks else None,
)
prefetched_chunks = []
# Kiro 重写后输出的是 Claude SSE 格式
# 客户端也是 Claude CLI不需要再进行格式转换
needs_conversion = False
ctx.needs_conversion = False
# 先处理预读的字节块
for chunk in prefetched_chunks:
buffer += chunk
# 处理缓冲区中的完整行
while b"\n" in buffer:
line_bytes, buffer = buffer.split(b"\n", 1)
try:
# 使用增量解码器,可以正确处理跨 chunk 的多字节字符
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
except Exception as e:
logger.warning(
f"[{self.request_id}] UTF-8 解码失败: {e}, "
f"bytes={line_bytes[:50]!r}"
)
continue
normalized_line = line.rstrip("\r")
events = sse_parser.feed_line(normalized_line)
if normalized_line == "":
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
self._mark_first_output(ctx, output_state)
yield b"\n"
continue
ctx.chunk_count += 1
# 格式转换或直接透传
if needs_conversion:
converted_lines, converted_events = self._convert_sse_line(
ctx, line, events
)
# 记录转换后的数据到 parsed_chunks
self._record_converted_chunks(ctx, converted_events)
for converted_line in converted_lines:
if converted_line:
self._mark_first_output(ctx, output_state)
yield (converted_line + "\n").encode("utf-8")
else:
self._mark_first_output(ctx, output_state)
yield (line + "\n").encode("utf-8")
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
if ctx.data_count > 0:
last_data_time = time.time()
# 继续处理剩余的流数据(使用同一个迭代器)
async for chunk in byte_iterator:
buffer += chunk
# 处理缓冲区中的完整行
while b"\n" in buffer:
line_bytes, buffer = buffer.split(b"\n", 1)
try:
# 使用增量解码器,可以正确处理跨 chunk 的多字节字符
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
except Exception as e:
logger.warning(
f"[{self.request_id}] UTF-8 解码失败: {e}, "
f"bytes={line_bytes[:50]!r}"
)
continue
normalized_line = line.rstrip("\r")
events = sse_parser.feed_line(normalized_line)
if normalized_line == "":
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
self._mark_first_output(ctx, output_state)
yield b"\n"
continue
ctx.chunk_count += 1
# 空流检测:超过阈值且无数据,发送错误事件并结束
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
elapsed = time.time() - last_data_time
if elapsed > self.DATA_TIMEOUT:
logger.warning("Provider '{}' 流超时且无数据", ctx.provider_name)
# 设置错误状态用于后续记录
ctx.status_code = 504
ctx.error_message = "流式响应超时,未收到有效数据"
ctx.upstream_response = f"流超时: Provider={ctx.provider_name}, elapsed={elapsed:.1f}s, chunk_count={ctx.chunk_count}, data_count=0"
error_event = {
"type": "error",
"error": {
"type": "empty_stream_timeout",
"message": ctx.error_message,
},
}
self._mark_first_output(ctx, output_state)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
return
# 格式转换或直接透传
if needs_conversion:
converted_lines, converted_events = self._convert_sse_line(
ctx, line, events
)
# 记录转换后的数据到 parsed_chunks
self._record_converted_chunks(ctx, converted_events)
for converted_line in converted_lines:
if converted_line:
self._mark_first_output(ctx, output_state)
yield (converted_line + "\n").encode("utf-8")
else:
self._mark_first_output(ctx, output_state)
yield (line + "\n").encode("utf-8")
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
if ctx.data_count > 0:
last_data_time = time.time()
# 处理剩余事件
flushed_events = sse_parser.flush()
for event in flushed_events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
# 检查是否收到数据
if ctx.data_count == 0:
# 空流通常意味着配置错误(如 base_url 指向了网页而非 API
logger.error(
f"Provider '{ctx.provider_name}' 返回空流式响应 (收到 {ctx.chunk_count} 个非数据行), "
f"可能是 endpoint base_url 配置错误"
)
# 设置错误状态用于后续记录
ctx.status_code = 503
ctx.error_message = "上游服务返回了空的流式响应"
ctx.upstream_response = f"空流式响应: Provider={ctx.provider_name}, chunk_count={ctx.chunk_count}, data_count=0, 可能是 base_url 配置错误"
error_event = {
"type": "error",
"error": {
"type": "empty_response",
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
logger.debug("流式数据转发完成")
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
client_fmt = (ctx.client_api_format or "").strip().lower()
if needs_conversion and client_fmt == "openai:chat":
yield b"data: [DONE]\n\n"
except GeneratorExit:
raise
except httpx.StreamClosed:
# 连接关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count == 0:
logger.warning("Provider '{}' 流连接关闭且无数据", ctx.provider_name)
# 设置错误状态用于后续记录
ctx.status_code = 503
ctx.error_message = "上游服务连接关闭且未返回数据"
ctx.upstream_response = f"流连接关闭: Provider={ctx.provider_name}, chunk_count={ctx.chunk_count}, data_count=0"
error_event = {
"type": "error",
"error": {
"type": "stream_closed",
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.RemoteProtocolError:
# 连接异常关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count > 0:
error_event = {
"type": "error",
"error": {
"type": "connection_error",
"message": "上游连接意外关闭,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
except httpx.ReadError:
# 代理/上游连接读取失败(如 aether-proxy 中断),与 RemoteProtocolError 处理逻辑一致
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count > 0:
error_event = {
"type": "error",
"error": {
"type": "connection_error",
"message": "代理或上游连接读取失败,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
finally:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass

View File

@@ -0,0 +1,319 @@
"""CLI Handler - 请求准备 Mixin"""
from __future__ import annotations
from typing import (
TYPE_CHECKING,
Any,
)
from src.api.handlers.base.utils import get_format_converter_registry
from src.core.logger import logger
if TYPE_CHECKING:
from src.core.api_format import EndpointDefinition
class CliRequestMixin:
"""请求准备相关方法的 Mixin"""
async def _get_mapped_model(
self,
source_model: str,
provider_id: str,
) -> str | None:
"""
获取模型映射后的实际模型名
查找逻辑:
1. 直接通过 GlobalModel.name 匹配
2. 查找该 Provider 的 Model 实现
3. 使用 provider_model_name / provider_model_mappings 选择最终名称
Args:
source_model: 用户请求的模型名(必须是 GlobalModel.name
provider_id: Provider ID
Returns:
映射后的 Provider 模型名,如果没有找到映射则返回 None
"""
from src.services.model.mapper import ModelMapperMiddleware
mapper = ModelMapperMiddleware(self.db)
mapping = await mapper.get_mapping(source_model, provider_id)
logger.debug(
f"[CLI] _get_mapped_model: source={source_model}, provider={provider_id[:8]}..., mapping={mapping}"
)
if mapping and mapping.model:
# 使用 select_provider_model_name 支持模型映射功能
# 传入 api_key.id 作为 affinity_key实现相同用户稳定选择同一映射
# 传入 api_format 用于过滤适用的映射作用域
affinity_key = self.api_key.id if self.api_key else None
mapped_name = mapping.model.select_provider_model_name(
affinity_key, api_format=self.FORMAT_ID
)
logger.debug(
f"[CLI] 模型映射: {source_model} -> {mapped_name} (provider={provider_id[:8]}...)"
)
return mapped_name
logger.debug("[CLI] 无模型映射,使用原始名称: {}", source_model)
return None
def extract_model_from_request(
self,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
) -> str:
"""
从请求中提取模型名 - 子类可覆盖
不同 API 格式的 model 位置不同:
- OpenAI/Claude: 在请求体中 request_body["model"]
- Gemini: 在 URL 路径中 path_params["model"]
子类应覆盖此方法实现各自的提取逻辑。
Args:
request_body: 请求体
path_params: URL 路径参数
Returns:
模型名,如果无法提取则返回 "unknown"
"""
# 默认实现:从请求体获取
model = request_body.get("model")
return str(model) if model else "unknown"
def apply_mapped_model(
self,
request_body: dict[str, Any],
mapped_model: str, # noqa: ARG002 - 子类使用
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
基类默认实现:不修改请求体,保持原样透传。
子类应覆盖此方法实现各自的模型名替换逻辑。
Args:
request_body: 原始请求体
mapped_model: 映射后的模型名(子类使用)
Returns:
请求体(默认不修改)
"""
# 基类不修改请求体,子类覆盖此方法实现特定格式的处理
return request_body
def prepare_provider_request_body(
self,
request_body: dict[str, Any],
) -> dict[str, Any]:
"""
准备发送给 Provider 的请求体 - 子类可覆盖
在模型映射之后、发送请求之前调用,用于移除不需要发送给上游的字段。
例如 Gemini API 需要移除请求体中的 model 字段(因为 model 在 URL 路径中)。
Args:
request_body: 经过模型映射处理后的请求体
Returns:
准备好的请求体
"""
return request_body
def finalize_provider_request(
self,
request_body: dict[str, Any],
*,
mapped_model: str | None,
provider_api_format: str | None,
) -> dict[str, Any]:
"""
格式转换完成后、envelope 之前的模型感知后处理钩子 - 子类可覆盖
用于根据目标模型的特性对请求体做最终调整,例如:
- 图像生成模型需要移除不兼容的 tools/system_instruction 并注入 imageConfig
- 特定模型需要注入/移除某些字段
- Gemini 格式:清理无效 parts 和合并连续同角色 contents
此方法在流式和非流式路径中均会被调用,且 mapped_model 已确定。
Args:
request_body: 已完成格式转换的请求体
mapped_model: 映射后的目标模型名
provider_api_format: Provider 侧 API 格式标识
Returns:
调整后的请求体
"""
# Gemini 格式请求:清理无效 parts 和合并连续同角色 contents
# 跨格式转换(如 Claude -> Gemini可能产生 thinking 等无法表示的块,
# 导致 parts 为空或缺少有效 data-oneof 字段,被 Google API 拒绝。
if provider_api_format and "gemini" in str(provider_api_format).lower():
contents = request_body.get("contents")
if isinstance(contents, list):
from src.core.api_format.conversion.normalizers.gemini import (
compact_gemini_contents,
)
request_body["contents"] = compact_gemini_contents(contents)
return request_body
@staticmethod
def _get_format_metadata(format_id: str) -> "EndpointDefinition | None":
"""获取 endpoint 元数据(解析失败返回 None"""
from src.core.api_format.metadata import resolve_endpoint_definition
return resolve_endpoint_definition(format_id)
def _finalize_converted_request(
self,
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
mapped_model: str | None,
fallback_model: str,
is_stream: bool,
) -> None:
"""
跨格式转换后统一设置并清理 model/stream 字段(原地修改)
处理逻辑:
1. 根据目标格式决定是否在 body 中设置 model
2. 若客户端格式不含 stream 字段但 Provider 需要,则显式设置
3. 移除目标格式不允许在 body 中携带的字段(如 Gemini 的 model/stream
Args:
request_body: 转换后的请求体(会被原地修改)
client_api_format: 客户端 API 格式
provider_api_format: Provider API 格式
mapped_model: 映射后的模型名
fallback_model: 备用模型名
is_stream: 是否流式请求
"""
client_meta = self._get_format_metadata(client_api_format)
provider_meta = self._get_format_metadata(provider_api_format)
# 默认model_in_body=True, stream_in_body=True如 OpenAI/Claude
client_uses_stream = client_meta.stream_in_body if client_meta else True
provider_model_in_body = provider_meta.model_in_body if provider_meta else True
provider_stream_in_body = provider_meta.stream_in_body if provider_meta else True
# 设置 model仅当 Provider 允许且 body 中需要)
if provider_model_in_body:
request_body["model"] = mapped_model or fallback_model
else:
request_body.pop("model", None)
# 设置 stream客户端不带但 Provider 需要时显式设置Provider 不需要时移除)
if provider_stream_in_body:
if not client_uses_stream:
request_body["stream"] = is_stream
else:
request_body.pop("stream", None)
# OpenAI Chat Completions: request usage in streaming mode.
provider_fmt = str(provider_api_format or "").strip().lower()
if is_stream and provider_fmt == "openai:chat":
stream_options = request_body.get("stream_options")
if not isinstance(stream_options, dict):
stream_options = {}
stream_options["include_usage"] = True
request_body["stream_options"] = stream_options
def _convert_request_for_cross_format(
self,
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
mapped_model: str | None,
fallback_model: str,
is_stream: bool,
*,
target_variant: str | None = None,
) -> tuple[dict[str, Any], str]:
"""
跨格式请求转换的公共逻辑
将客户端格式的请求体转换为 Provider 格式,并处理 model/stream 字段的补齐和清理。
Args:
request_body: 原始请求体(会被修改)
client_api_format: 客户端 API 格式
provider_api_format: Provider API 格式
mapped_model: 映射后的模型名
fallback_model: 备用模型名(通常是原始请求的 model
is_stream: 是否流式请求
target_variant: 目标变体(如 "codex"),用于同格式但有细微差异的上游
Returns:
(转换后的请求体, 用于 URL 的模型名)
"""
registry = get_format_converter_registry()
converted_body = registry.convert_request(
request_body,
str(client_api_format),
str(provider_api_format),
target_variant=target_variant,
)
# 先计算 URL 模型(在清理 body 中的 model 字段之前)
url_model = (
self.get_model_for_url(converted_body, mapped_model) or mapped_model or fallback_model
)
# 统一设置并清理 model/stream 字段
self._finalize_converted_request(
converted_body,
str(client_api_format),
str(provider_api_format),
mapped_model,
fallback_model,
is_stream,
)
return converted_body, url_model
def get_model_for_url(
self,
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None:
"""
获取用于 URL 路径的模型名
某些 API 格式(如 Gemini需要将 model 放入 URL 路径中。
子类应覆盖此方法返回正确的值。
Args:
request_body: 请求体
mapped_model: 映射后的模型名(如果有)
Returns:
用于 URL 路径的模型名,默认优先使用映射后的名称
"""
return mapped_model or request_body.get("model")
def _extract_response_metadata(
self,
response: dict[str, Any],
) -> dict[str, Any]:
"""
从响应中提取 Provider 特有的元数据 - 子类可覆盖
例如 Gemini 返回的 modelVersion 字段。
这些元数据会存储到 Usage.request_metadata 中。
Args:
response: Provider 返回的响应
Returns:
元数据字典,默认为空
"""
return {}

View File

@@ -0,0 +1,105 @@
"""SSE 解析辅助函数"""
from __future__ import annotations
import json
from typing import Any
from src.core.logger import logger
def _parse_sse_data_line(line: str) -> tuple[Any | None, str]:
"""
解析标准 SSE data 行
Args:
line: 以 "data:" 开头的 SSE 行
Returns:
(parsed_json, status) 元组:
- (parsed_dict, "ok") - 解析成功
- (None, "empty") - 内容为空
- (None, "invalid") - JSON 解析失败,调用方应透传原始行
"""
data_content = line[5:].strip()
if not data_content:
return None, "empty"
try:
return json.loads(data_content), "ok"
except json.JSONDecodeError:
return None, "invalid"
def _parse_sse_event_data_line(line: str) -> tuple[Any | None, str]:
"""
解析 event + data 同行格式(如 "event: xxx data: {...}"
Args:
line: 以 "event:" 开头且包含 " data:" 的 SSE 行
Returns:
(parsed_json, status) 元组
"""
_event_part, data_part = line.split(" data:", 1)
data_content = data_part.strip()
try:
return json.loads(data_content), "ok"
except json.JSONDecodeError:
return None, "invalid"
def _parse_gemini_json_array_line(line: str) -> tuple[Any | None, str]:
"""
解析 Gemini JSON-array 格式的裸 JSON 行
Gemini 流式响应可能是 JSON 数组格式,每行是数组元素。
Args:
line: 原始行(可能是 "[", "]", ",", 或 JSON 对象)
Returns:
(parsed_json, status) 元组
"""
stripped = line.strip()
if stripped in ("", "[", "]", ","):
return None, "skip"
candidate = stripped.lstrip(",").rstrip(",").strip()
try:
return json.loads(candidate), "ok"
except json.JSONDecodeError:
logger.debug("Gemini JSON-array line skip: {}", stripped[:50])
return None, "invalid"
def _format_converted_events_to_sse(
converted_events: list[dict[str, Any]],
client_format: str,
) -> list[str]:
"""
将转换后的事件格式化为 SSE 行
Args:
converted_events: 转换后的事件列表
client_format: 客户端 API 格式
Returns:
SSE 行列表(每个元素是完整的 SSE 事件,包含尾部空行)
"""
result: list[str] = []
needs_event_line = str(client_format or "").strip().lower().startswith("claude:")
for evt in converted_events:
payload = json.dumps(evt, ensure_ascii=False)
if needs_event_line:
evt_type = evt.get("type") if isinstance(evt, dict) else None
if isinstance(evt_type, str) and evt_type:
# Claude 格式event + data + 空行
result.append(f"event: {evt_type}\ndata: {payload}\n")
else:
result.append(f"data: {payload}\n")
else:
# OpenAI 格式data + 空行
result.append(f"data: {payload}\n")
return result

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,650 @@
"""CLI Handler - 同步处理 Mixin"""
from __future__ import annotations
import json
import time
from typing import TYPE_CHECKING, Any
import httpx
from fastapi.responses import JSONResponse
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.request_builder import get_provider_auth
from src.api.handlers.base.stream_context import extract_proxy_timing, is_format_converted
from src.api.handlers.base.upstream_stream_bridge import (
aggregate_upstream_stream_to_internal_response,
)
from src.api.handlers.base.utils import (
filter_proxy_response_headers,
get_format_converter_registry,
)
from src.config.settings import config
from src.core.error_utils import extract_client_error_message
from src.core.exceptions import (
ProviderAuthException,
ProviderNotAvailableException,
ProviderRateLimitException,
ProviderTimeoutException,
ThinkingSignatureException,
)
from src.core.logger import logger
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.provider.behavior import get_provider_behavior
from src.services.provider.stream_policy import (
enforce_stream_mode_for_upstream,
get_upstream_stream_policy,
resolve_upstream_is_stream,
)
from src.services.provider.transport import build_provider_url
if TYPE_CHECKING:
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
class CliSyncMixin:
"""同步处理相关方法的 Mixin"""
async def process_sync(
self,
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""
处理非流式请求
通用流程:
1. 构建请求
2. 通过 TaskService/FailoverEngine 执行
3. 解析响应并记录统计
"""
logger.debug("开始非流式响应处理 ({})", self.FORMAT_ID)
# 使用子类实现的方法提取 model不同 API 格式的 model 位置不同)
model = self.extract_model_from_request(original_request_body, path_params)
api_format = self.allowed_api_formats[0]
sync_start_time = time.time()
# 提前创建 pending 记录,让前端可以立即看到"处理中"
self._create_pending_usage(
model=model,
is_stream=False,
request_type="chat",
api_format=self.FORMAT_ID,
request_headers=original_headers,
request_body=original_request_body,
)
provider_name = None
response_json = None
status_code = 200
response_headers = {}
provider_api_format = "" # 用于追踪 Provider 的 API 格式
provider_request_headers = {} # 发送给 Provider 的请求头
provider_request_body = None # 实际发送给 Provider 的请求体
provider_id = None # Provider ID用于失败记录
endpoint_id = None # Endpoint ID用于失败记录
key_id = None # Key ID用于失败记录
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
response_metadata_result: dict[str, Any] = {} # Provider 响应元数据
needs_conversion = False # 是否需要格式转换(由 candidate 决定)
sync_proxy_info: dict[str, Any] | None = None # 代理信息
# 可变请求体容器:允许 TaskService 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: dict[str, Any] = {"body": original_request_body}
async def sync_request_func(
provider: "Provider",
endpoint: "ProviderEndpoint",
key: "ProviderAPIKey",
candidate: ProviderCandidate,
) -> dict[str, Any]:
nonlocal provider_name, response_json, status_code, response_headers, provider_api_format, provider_request_headers, provider_request_body, mapped_model_result, response_metadata_result, needs_conversion, sync_proxy_info
provider_name = str(provider.name)
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
# 获取模型映射(优先使用映射匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=model,
provider_id=str(provider.id),
)
# 应用模型映射到请求体(子类可覆盖此方法处理不同格式)
if mapped_model:
mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录
request_body = self.apply_mapped_model(request_body_ref["body"], mapped_model)
else:
request_body = dict(request_body_ref["body"])
client_api_format = (
api_format.value if hasattr(api_format, "value") else str(api_format)
)
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
behavior = get_provider_behavior(
provider_type=provider_type,
endpoint_sig=provider_api_format,
)
envelope = behavior.envelope
target_variant = behavior.same_format_variant
# 跨格式转换也允许变体Antigravity 需要保留/翻译 Claude thinking 块)
conversion_variant = behavior.cross_format_variant
# Upstream streaming policy (per-endpoint).
upstream_policy = get_upstream_stream_policy(
endpoint,
provider_type=provider_type,
endpoint_sig=provider_api_format,
)
upstream_is_stream = resolve_upstream_is_stream(
client_is_stream=False,
policy=upstream_policy,
)
# 跨格式:先做请求体转换(失败触发 failover
if needs_conversion and provider_api_format:
request_body, url_model = self._convert_request_for_cross_format(
request_body,
client_api_format,
provider_api_format,
mapped_model,
model,
is_stream=upstream_is_stream,
target_variant=conversion_variant,
)
else:
# 同格式:按原逻辑做轻量清理(子类可覆盖)
request_body = self.prepare_provider_request_body(request_body)
url_model = (
self.get_model_for_url(request_body, mapped_model) or mapped_model or model
)
# 同格式时也需要应用 target_variant 转换(如 Codex
if target_variant and provider_api_format:
registry = get_format_converter_registry()
request_body = registry.convert_request(
request_body,
provider_api_format,
provider_api_format,
target_variant=target_variant,
)
# 模型感知的请求后处理(如图像生成模型移除不兼容字段)
request_body = self.finalize_provider_request(
request_body,
mapped_model=mapped_model,
provider_api_format=provider_api_format,
)
# Force upstream stream/sync mode in request body (best-effort).
if provider_api_format:
enforce_stream_mode_for_upstream(
request_body,
provider_api_format=provider_api_format,
upstream_is_stream=upstream_is_stream,
)
# 获取认证信息(处理 Service Account 等异步认证场景)
auth_info = await get_provider_auth(endpoint, key)
# Provider envelope: wrap request after auth is available and before RequestBuilder.build().
if envelope:
request_body, url_model = envelope.wrap_request(
request_body,
model=url_model or model or "",
url_model=url_model,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
# Provider envelope: extra upstream headers (e.g. dedicated User-Agent).
extra_headers: dict[str, str] = {}
if envelope:
extra_headers.update(envelope.extra_headers() or {})
# 使用 RequestBuilder 构建请求体和请求头
# 注意mapped_model 已经应用到 request_body这里不再传递
# 上游始终使用 header 认证,不跟随客户端的 query 方式
provider_payload, provider_headers = self._request_builder.build(
request_body,
original_headers,
endpoint,
key,
is_stream=upstream_is_stream,
extra_headers=extra_headers if extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
)
if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent
set_accept_if_absent(provider_headers)
# 保存发送给 Provider 的请求信息(用于调试和统计)
provider_request_headers = provider_headers
provider_request_body = provider_payload
url = build_provider_url(
endpoint,
query_params=query_params,
path_params={"model": url_model},
is_stream=upstream_is_stream, # sync handler may still force upstream streaming
key=key,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
# 非流式:必须在 build_provider_url 调用后立即缓存(避免 contextvar 被后续调用覆盖)
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
# 解析有效代理Key 级别优先于 Provider 级别)
from src.services.proxy_node.resolver import (
get_proxy_label,
resolve_effective_proxy,
resolve_proxy_info,
)
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
sync_proxy_info = resolve_proxy_info(_effective_proxy)
_proxy_label = get_proxy_label(sync_proxy_info)
logger.info(
f" └─ [{self.request_id}] 发送{'上游流式(聚合)' if upstream_is_stream else '非流式'}请求: "
f"Provider={provider.name}, Endpoint={endpoint.id[:8] if endpoint.id else 'N/A'}..., "
f"Key=***{key.api_key[-4:] if key.api_key else 'N/A'}, "
f"原始模型={model}, 映射后={mapped_model or '无映射'}, URL模型={url_model}, "
f"代理={_proxy_label}"
)
# 获取复用的 HTTP 客户端支持代理配置Key 级别优先于 Provider 级别)
# 注意:使用 get_proxy_client 复用连接池,不再每次创建新客户端
from src.clients.http_client import HTTPClientPool
from src.services.proxy_node.resolver import (
build_post_kwargs,
build_stream_kwargs,
resolve_delegate_config,
)
# 非流式请求使用 http_request_timeout 作为整体超时
# 优先使用 Provider 配置,否则使用全局配置
request_timeout = provider.request_timeout or config.http_request_timeout
delegate_cfg = resolve_delegate_config(_effective_proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=_effective_proxy
)
# 注意:不使用 async with因为复用的客户端不应该被关闭
# 超时通过 timeout 参数控制
resp: httpx.Response | None = None
if not upstream_is_stream:
try:
_pkw = build_post_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
payload=provider_payload,
timeout=request_timeout,
)
resp = await http_client.post(**_pkw)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
if selected_base_url_cached:
logger.warning(
f"[{envelope.name}] Connection error: {selected_base_url_cached} ({e})"
)
raise
else:
# Forced upstream streaming: aggregate SSE to a sync JSON response.
registry = get_format_converter_registry()
provider_parser = (
get_parser_for_format(provider_api_format) if provider_api_format else None
)
try:
_stream_args = build_stream_kwargs(
delegate_cfg,
url=url,
headers=provider_headers,
payload=provider_payload,
timeout=request_timeout,
)
async with http_client.stream(**_stream_args) as stream_resp:
resp = stream_resp
status_code = stream_resp.status_code
response_headers = dict(stream_resp.headers)
extract_proxy_timing(sync_proxy_info, response_headers)
if envelope:
envelope.on_http_status(
base_url=selected_base_url_cached,
status_code=status_code,
)
stream_resp.raise_for_status()
byte_iter = stream_resp.aiter_bytes()
if provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
from src.services.provider.adapters.kiro.eventstream_rewriter import (
apply_kiro_stream_rewrite,
)
byte_iter = apply_kiro_stream_rewrite(byte_iter, model=str(model or ""))
internal_resp = await aggregate_upstream_stream_to_internal_response(
byte_iter,
provider_api_format=provider_api_format,
provider_name=str(provider.name),
model=str(model or ""),
request_id=str(self.request_id or ""),
envelope=envelope,
provider_parser=provider_parser,
)
tgt_norm = (
registry.get_normalizer(client_api_format)
if client_api_format
else None
)
if tgt_norm is None:
raise RuntimeError(f"未注册 Normalizer: {client_api_format}")
response_json = tgt_norm.response_from_internal(
internal_resp,
requested_model=model,
)
response_json = response_json if isinstance(response_json, dict) else {}
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
if selected_base_url_cached:
logger.warning(
f"[{envelope.name}] Connection error: {selected_base_url_cached} ({e})"
)
raise
status_code = resp.status_code
response_headers = dict(resp.headers)
extract_proxy_timing(sync_proxy_info, response_headers)
if envelope:
envelope.on_http_status(base_url=selected_base_url_cached, status_code=status_code)
# Forced upstream streaming already built response_json via aggregator.
if upstream_is_stream:
response_metadata_result = self._extract_response_metadata(response_json or {})
return response_json if isinstance(response_json, dict) else {}
# Reuse HTTPStatusError classification path (handled by TaskService/error_classifier).
try:
resp.raise_for_status()
except httpx.HTTPStatusError as e:
error_body = ""
try:
error_body = resp.text[:4000] if resp.text else ""
except Exception:
error_body = ""
e.upstream_response = error_body # type: ignore[attr-defined]
raise
# 安全解析 JSON 响应,处理可能的编码错误
try:
response_json = resp.json()
except (UnicodeDecodeError, json.JSONDecodeError) as e:
# 获取原始响应内容用于调试(存入 upstream_response
content_type = resp.headers.get("content-type", "unknown")
content_encoding = resp.headers.get("content-encoding", "none")
raw_content = ""
try:
raw_content = resp.text[:500] if resp.text else "(empty)"
except Exception:
try:
raw_content = repr(resp.content[:500]) if resp.content else "(empty)"
except Exception:
raw_content = "(unable to read)"
logger.error(
f"[{self.request_id}] 无法解析响应 JSON: {e}, "
f"Content-Type: {content_type}, Content-Encoding: {content_encoding}, "
f"响应长度: {len(resp.content)} bytes, 原始内容: {raw_content}"
)
# 判断错误类型,生成友好的客户端错误消息(不暴露提供商信息)
if raw_content == "(empty)" or not raw_content.strip():
client_message = "上游服务返回了空响应"
elif raw_content.strip().startswith(("<", "<!doctype", "<!DOCTYPE")):
client_message = "上游服务返回了非预期的响应格式"
else:
client_message = "上游服务返回了无效的响应"
raise ProviderNotAvailableException(
client_message,
provider_name=str(provider.name),
upstream_status=resp.status_code,
upstream_response=raw_content,
)
if envelope:
response_json = envelope.unwrap_response(response_json)
envelope.postprocess_unwrapped_response(model=model, data=response_json)
# 提取 Provider 响应元数据(子类可覆盖)
response_metadata_result = self._extract_response_metadata(response_json)
return response_json if isinstance(response_json, dict) else {}
try:
# 解析能力需求
capability_requirements = self._resolve_capability_requirements(
model_name=model,
request_headers=original_headers,
request_body=original_request_body,
)
preferred_key_ids = await self._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
# 统一入口:总是通过 TaskService
from src.services.task import TaskService
from src.services.task.context import TaskMode
exec_result = await TaskService(self.db, self.redis).execute(
task_type="cli",
task_mode=TaskMode.SYNC,
api_format=api_format,
model_name=model,
user_api_key=self.api_key,
request_func=sync_request_func,
request_id=self.request_id,
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref,
)
result = exec_result.response
actual_provider_name = exec_result.provider_name or "unknown"
attempt_id = exec_result.request_candidate_id
provider_id = exec_result.provider_id
endpoint_id = exec_result.endpoint_id
key_id = exec_result.key_id
provider_name = actual_provider_name
response_time_ms = int((time.time() - sync_start_time) * 1000)
# 确保 response_json 不为 None
if response_json is None:
response_json = {}
# 跨格式:响应转换回 client_format失败不触发 failover保守回退为原始响应
if (
needs_conversion
and provider_api_format
and api_format
and isinstance(response_json, dict)
):
try:
registry = get_format_converter_registry()
response_json = registry.convert_response(
response_json,
provider_api_format,
api_format,
requested_model=model, # 使用用户请求的原始模型名
)
logger.debug(
"非流式响应格式转换完成: {} -> {}", provider_api_format, api_format
)
except Exception as conv_err:
logger.warning("非流式响应格式转换失败,使用原始响应: {}", conv_err)
# 使用解析器提取 usage
usage = self.parser.extract_usage_from_response(response_json)
input_tokens = usage.get("input_tokens", 0)
output_tokens = usage.get("output_tokens", 0)
cached_tokens = usage.get("cache_read_tokens", 0)
cache_creation_tokens = usage.get("cache_creation_tokens", 0)
output_text = self.parser.extract_text_content(response_json)[:200]
# 使用实际发送给 Provider 的请求体(如果有),否则用原始请求体
actual_request_body = provider_request_body or original_request_body
# 非流式成功时,返回给客户端的是提供商响应头(透传)
client_response_headers = filter_proxy_response_headers(response_headers)
client_response_headers["content-type"] = "application/json"
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
total_cost = await self.telemetry.record_success(
provider=provider_name,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
response_time_ms=response_time_ms,
status_code=status_code,
request_headers=original_headers,
request_body=actual_request_body,
response_headers=response_headers,
client_response_headers=client_response_headers,
response_body=response_json,
cache_creation_tokens=cache_creation_tokens,
cache_read_tokens=cached_tokens,
is_stream=False,
provider_request_headers=provider_request_headers,
api_format=api_format,
# 格式转换追踪
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, str(api_format)),
# Provider 侧追踪信息(用于记录真实成本)
provider_id=provider_id,
provider_endpoint_id=endpoint_id,
provider_api_key_id=key_id,
# 模型映射信息
target_model=mapped_model_result,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata=response_metadata_result if response_metadata_result else None,
request_metadata=request_metadata or None,
)
logger.info("{} 非流式响应处理完成", self.FORMAT_ID)
# 透传提供商的响应头
return JSONResponse(
status_code=status_code,
content=response_json,
headers=client_response_headers,
)
except ThinkingSignatureException as e:
# Thinking 签名错误TaskService 层已处理整流重试但仍失败
# 记录实际发送给 Provider 的请求体,便于排查问题根因
response_time_ms = int((time.time() - sync_start_time) * 1000)
actual_request_body = provider_request_body or original_request_body
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=e.status_code or 400,
request_headers=original_headers,
request_body=actual_request_body,
error_message=str(e),
is_stream=False,
api_format=api_format,
request_metadata=request_metadata or None,
)
raise
except Exception as e:
response_time_ms = int((time.time() - sync_start_time) * 1000)
status_code = 503
if isinstance(e, ProviderAuthException):
status_code = 503
elif isinstance(e, ProviderRateLimitException):
status_code = 429
elif isinstance(e, ProviderTimeoutException):
status_code = 504
# 使用实际发送给 Provider 的请求体(如果有),否则用原始请求体
actual_request_body = provider_request_body or original_request_body
# 尝试从异常中提取响应头
error_response_headers: dict[str, str] = {}
if isinstance(e, ProviderRateLimitException) and e.response_headers:
error_response_headers = e.response_headers
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
error_response_headers = dict(e.response.headers)
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=status_code,
error_message=extract_client_error_message(e),
request_headers=original_headers,
request_body=actual_request_body,
is_stream=False,
api_format=api_format,
provider_request_headers=provider_request_headers,
response_headers=error_response_headers,
# 非流式失败返回给客户端的是 JSON 错误响应
client_response_headers={"content-type": "application/json"},
# 格式转换追踪
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, str(api_format)),
# 模型映射信息
target_model=mapped_model_result,
request_metadata=request_metadata or None,
)
raise
async def _extract_error_text(self, e: httpx.HTTPStatusError) -> str:
"""从 HTTP 错误中提取错误文本"""
try:
if hasattr(e.response, "is_stream_consumed") and not e.response.is_stream_consumed:
error_bytes = await e.response.aread()
for encoding in ["utf-8", "gbk", "latin1"]:
try:
return error_bytes.decode(encoding)
except (UnicodeDecodeError, LookupError):
continue
return error_bytes.decode("utf-8", errors="replace")
else:
return (
e.response.text
if hasattr(e.response, "_content")
else "Unable to read response"
)
except Exception as decode_error:
return f"Unable to read error response: {decode_error}"

View File

@@ -0,0 +1,10 @@
"""格式转换层常量定义。
将跨层共享的常量集中在 core 层,避免 core -> services 的反向依赖。
"""
from __future__ import annotations
# Thinking 签名验证的跳过标记
# 当无法获取真实签名时,使用此值作为占位符
DUMMY_THOUGHT_SIGNATURE = "skip_thought_signature_validator"

View File

@@ -379,14 +379,11 @@ class GeminiNormalizer(FormatNormalizer):
isinstance(p, dict) and p.get("thought") is True for p in parts
)
if not has_thought:
try:
from src.services.provider.adapters.antigravity.constants import (
DUMMY_THOUGHT_SIGNATURE,
)
from src.core.api_format.conversion.constants import (
DUMMY_THOUGHT_SIGNATURE,
)
dummy_sig = DUMMY_THOUGHT_SIGNATURE
except Exception:
dummy_sig = "skip_thought_signature_validator"
dummy_sig = DUMMY_THOUGHT_SIGNATURE
dummy_part: dict[str, Any] = {
"text": "Thinking...",
@@ -1423,10 +1420,11 @@ class GeminiNormalizer(FormatNormalizer):
if signature is None and target_variant == "antigravity":
model_str = str(model or "")
from src.core.api_format.conversion.constants import (
DUMMY_THOUGHT_SIGNATURE,
)
try:
from src.services.provider.adapters.antigravity.constants import (
DUMMY_THOUGHT_SIGNATURE,
)
from src.services.provider.adapters.antigravity.signature_cache import (
signature_cache,
)
@@ -1445,7 +1443,7 @@ class GeminiNormalizer(FormatNormalizer):
except Exception:
# Best-effort fallback: Gemini models can accept a dummy signature.
if model_str.startswith("gemini-"):
signature = "skip_thought_signature_validator"
signature = DUMMY_THOUGHT_SIGNATURE
# For Antigravity, missing signature is likely to fail upstream validation.
if target_variant == "antigravity" and not signature:
@@ -1564,8 +1562,9 @@ class GeminiNormalizer(FormatNormalizer):
payload_sig = None
signature: str | None = None
from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE
try:
from src.services.provider.adapters.antigravity.constants import DUMMY_THOUGHT_SIGNATURE
from src.services.provider.adapters.antigravity.signature_cache import signature_cache
cached_or_dummy = signature_cache.get_or_dummy(model, text_val)

10
src/models/_base.py Normal file
View File

@@ -0,0 +1,10 @@
"""
SQLAlchemy Base 声明基类
所有数据库模型子模块从此文件导入 Base确保全局唯一。
直接复用 database.py 的 Base避免出现两套 MetaData 实例。
"""
from src.models.database import Base
__all__ = ["Base"]

178
src/models/auth.py Normal file
View File

@@ -0,0 +1,178 @@
"""
认证相关数据库模型
包含: LDAPConfig, OAuthProvider, UserOAuthLink
"""
from __future__ import annotations
import uuid
from datetime import datetime, timezone
from sqlalchemy import (
JSON,
Boolean,
Column,
DateTime,
ForeignKey,
Integer,
String,
Text,
UniqueConstraint,
)
from sqlalchemy.orm import relationship
from ._base import Base
class LDAPConfig(Base):
"""LDAP认证配置表 - 单行配置"""
__tablename__ = "ldap_configs"
id = Column(Integer, primary_key=True, autoincrement=True)
server_url = Column(String(255), nullable=False) # ldap://host:389 或 ldaps://host:636
bind_dn = Column(Text, nullable=False) # 绑定账号 DN可能很长
bind_password_encrypted = Column(Text, nullable=True) # 加密的绑定密码(允许 NULL 表示已清除)
base_dn = Column(Text, nullable=False) # 用户搜索基础 DN可能很长
user_search_filter = Column(
Text, default="(uid={username})", nullable=False
) # 用户搜索过滤器(可能很复杂)
username_attr = Column(
String(50), default="uid", nullable=False
) # 用户名属性 (uid/sAMAccountName)
email_attr = Column(String(50), default="mail", nullable=False) # 邮箱属性
display_name_attr = Column(String(50), default="cn", nullable=False) # 显示名称属性
is_enabled = Column(Boolean, default=False, nullable=False) # 是否启用 LDAP 认证
is_exclusive = Column(
Boolean, default=False, nullable=False
) # 是否仅允许 LDAP 登录(禁用本地认证)
use_starttls = Column(Boolean, default=False, nullable=False) # 是否使用 STARTTLS
connect_timeout = Column(Integer, default=10, nullable=False) # 连接超时时间(秒)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
def set_bind_password(self, password: str) -> None:
"""
设置并加密绑定密码
Args:
password: 明文密码
"""
from src.core.crypto import crypto_service
self.bind_password_encrypted = crypto_service.encrypt(password)
def get_bind_password(self) -> str:
"""
获取解密后的绑定密码
Returns:
str: 解密后的明文密码
Raises:
DecryptionException: 解密失败时抛出异常
"""
from src.core.crypto import crypto_service
if not self.bind_password_encrypted:
return ""
return crypto_service.decrypt(self.bind_password_encrypted)
class OAuthProvider(Base):
"""OAuth Provider 配置表(按 provider_type 唯一)"""
__tablename__ = "oauth_providers"
# 使用 provider_type 作为主键,便于通过 URL 参数直接定位配置
provider_type = Column(String(50), primary_key=True)
display_name = Column(String(100), nullable=False)
client_id = Column(Text, nullable=False) # 某些 OAuth 提供商可能使用很长的 client_id
client_secret_encrypted = Column(Text, nullable=True) # 允许 NULL 表示尚未配置/已清除
# 可选覆盖端点(需在业务层做白名单校验)
authorization_url_override = Column(String(500), nullable=True)
token_url_override = Column(String(500), nullable=True)
userinfo_url_override = Column(String(500), nullable=True)
# 可选覆盖 scopesJSON 列表)
scopes = Column(JSON, nullable=True)
# 服务端控制 redirect_uri 与前端回调 URL
redirect_uri = Column(String(500), nullable=False)
frontend_callback_url = Column(String(500), nullable=False)
# Provider 特定配置/映射
attribute_mapping = Column(JSON, nullable=True)
extra_config = Column(JSON, nullable=True)
is_enabled = Column(Boolean, default=False, nullable=False)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
def set_client_secret(self, secret: str) -> None:
"""设置并加密 client_secret"""
from src.core.crypto import crypto_service
self.client_secret_encrypted = crypto_service.encrypt(secret)
def get_client_secret(self) -> str:
"""获取解密后的 client_secret未配置时返回空串"""
from src.core.crypto import crypto_service
if not self.client_secret_encrypted:
return ""
return crypto_service.decrypt(self.client_secret_encrypted)
class UserOAuthLink(Base):
"""用户与 OAuth Provider 的绑定关系"""
__tablename__ = "user_oauth_links"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
user_id = Column(
String(36),
ForeignKey("users.id", ondelete="CASCADE"),
nullable=False,
index=True,
)
provider_type = Column(
String(50),
ForeignKey("oauth_providers.provider_type", ondelete="CASCADE"),
nullable=False,
index=True,
)
provider_user_id = Column(String(255), nullable=False)
provider_username = Column(String(255), nullable=True)
provider_email = Column(String(255), nullable=True)
extra_data = Column(JSON, nullable=True)
linked_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
last_login_at = Column(DateTime(timezone=True), nullable=True)
__table_args__ = (
UniqueConstraint("provider_type", "provider_user_id", name="uq_oauth_provider_user"),
UniqueConstraint("user_id", "provider_type", name="uq_user_oauth_provider"),
)

360
src/models/misc.py Normal file
View File

@@ -0,0 +1,360 @@
"""
杂项数据库模型
包含: SystemConfig, VideoTask, Announcement, AnnouncementRead, AuditEventType, AuditLog,
GeminiFileMapping, _generate_short_id
"""
from __future__ import annotations
import secrets
import string
import uuid
from datetime import datetime, timezone
from enum import Enum as PyEnum
from sqlalchemy import (
JSON,
BigInteger,
Boolean,
Column,
DateTime,
Float,
ForeignKey,
Index,
Integer,
String,
Text,
UniqueConstraint,
)
from sqlalchemy.orm import relationship
from ._base import Base
class SystemConfig(Base):
"""系统配置表"""
__tablename__ = "system_configs"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
key = Column(String(100), unique=True, nullable=False)
value = Column(JSON, nullable=False)
description = Column(Text, nullable=True)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
def _generate_short_id(length: int = 12) -> str:
"""生成 Gemini 风格的短 ID小写字母+数字)"""
alphabet = string.ascii_lowercase + string.digits
return "".join(secrets.choice(alphabet) for _ in range(length))
class VideoTask(Base):
"""视频生成任务"""
__tablename__ = "video_tasks"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# Gemini 风格的短 ID用于对外暴露如 operations/xxx
short_id = Column(String(16), unique=True, index=True, default=_generate_short_id)
request_id = Column(
String(100), unique=True, index=True, nullable=False
) # 关联 Usage/RequestCandidate
external_task_id = Column(String(200))
# 关联
user_id = Column(String(36), ForeignKey("users.id"), nullable=False)
api_key_id = Column(String(36), ForeignKey("api_keys.id"))
provider_id = Column(String(36), ForeignKey("providers.id"))
endpoint_id = Column(String(36), ForeignKey("provider_endpoints.id"))
key_id = Column(String(36), ForeignKey("provider_api_keys.id"))
# 格式转换追踪
client_api_format = Column(String(50), nullable=False)
provider_api_format = Column(String(50), nullable=False)
format_converted = Column(Boolean, default=False)
# 任务配置
model = Column(String(100), nullable=False)
prompt = Column(Text, nullable=False)
original_request_body = Column(JSON)
converted_request_body = Column(JSON)
# 视频参数 (统一内部格式)
duration_seconds = Column(Integer, default=4)
resolution = Column(String(20), default="720p")
aspect_ratio = Column(String(10), default="16:9")
size = Column(String(20))
# 状态
status = Column(String(20), default="pending")
progress_percent = Column(Integer, default=0)
progress_message = Column(String(500))
# 结果
video_url = Column(String(2000))
video_urls = Column(JSON)
thumbnail_url = Column(String(2000))
video_size_bytes = Column(BigInteger)
video_duration_seconds = Column(Float) # 实际视频时长(秒)
video_expires_at = Column(DateTime(timezone=True))
# 存储 (可选)
stored_video_path = Column(String(500))
storage_provider = Column(String(50))
# 错误
error_code = Column(String(50))
error_message = Column(Text)
retry_count = Column(Integer, default=0)
max_retries = Column(Integer, default=3)
# 轮询配置
poll_interval_seconds = Column(Integer, default=10)
next_poll_at = Column(DateTime(timezone=True)) # 索引在 __table_args__ 中定义
poll_count = Column(Integer, default=0)
max_poll_count = Column(Integer, default=360)
# Remix 支持
remixed_from_task_id = Column(
String(36), ForeignKey("video_tasks.id", ondelete="SET NULL"), nullable=True
)
# 使用追踪(候选 key、请求头等
request_metadata = Column(JSON, nullable=True) # 存储候选 key 列表、请求头等追踪信息
# 示例: {
# "candidate_keys": [{"key_id": "xxx", "endpoint_id": "yyy", "priority": 1}, ...],
# "selected_key_index": 0,
# "client_ip": "1.2.3.4",
# "user_agent": "...",
# "request_headers": {...}
# }
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
submitted_at = Column(DateTime(timezone=True))
completed_at = Column(DateTime(timezone=True))
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
user = relationship("User", backref="video_tasks")
remixed_from = relationship("VideoTask", remote_side=[id], backref="remixes")
# 复合索引和唯一约束
__table_args__ = (
Index("idx_video_tasks_user_status", "user_id", "status"),
Index("idx_video_tasks_next_poll", "next_poll_at"),
Index("idx_video_tasks_external_id", "external_task_id"),
UniqueConstraint("user_id", "external_task_id", name="uq_video_tasks_user_external_id"),
)
class Announcement(Base):
"""公告表"""
__tablename__ = "announcements"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
title = Column(String(200), nullable=False)
content = Column(Text, nullable=False) # 支持 Markdown
type = Column(String(20), default="info") # info, warning, maintenance, important
priority = Column(Integer, default=0) # 优先级,数字越大越重要
# 发布信息
author_id = Column(String(36), ForeignKey("users.id", ondelete="SET NULL"), nullable=True)
is_active = Column(Boolean, default=True, index=True)
is_pinned = Column(Boolean, default=False) # 置顶
# 时间范围
start_time = Column(DateTime(timezone=True), nullable=True)
end_time = Column(DateTime(timezone=True), nullable=True)
# 时间戳
created_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
nullable=False,
index=True,
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
author = relationship("User", back_populates="authored_announcements")
reads = relationship(
"AnnouncementRead", back_populates="announcement", cascade="all, delete-orphan"
)
class AnnouncementRead(Base):
"""公告已读记录表"""
__tablename__ = "announcement_reads"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
user_id = Column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False)
announcement_id = Column(String(36), ForeignKey("announcements.id"), nullable=False)
read_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
# 唯一约束
__table_args__ = (UniqueConstraint("user_id", "announcement_id", name="uq_user_announcement"),)
# 关系
user = relationship("User", back_populates="announcement_reads")
announcement = relationship("Announcement", back_populates="reads")
class AuditEventType(PyEnum):
"""审计事件类型"""
# 认证相关
LOGIN_SUCCESS = "login_success"
LOGIN_FAILED = "login_failed"
LOGOUT = "logout"
API_KEY_CREATED = "api_key_created"
API_KEY_DELETED = "api_key_deleted"
API_KEY_USED = "api_key_used"
# 请求相关
REQUEST_SUCCESS = "request_success"
REQUEST_FAILED = "request_failed"
REQUEST_RATE_LIMITED = "request_rate_limited"
REQUEST_QUOTA_EXCEEDED = "request_quota_exceeded"
# 管理操作
USER_CREATED = "user_created"
USER_UPDATED = "user_updated"
USER_DELETED = "user_deleted"
PROVIDER_ADDED = "provider_added"
PROVIDER_UPDATED = "provider_updated"
PROVIDER_REMOVED = "provider_removed"
# 安全事件
SUSPICIOUS_ACTIVITY = "suspicious_activity"
UNAUTHORIZED_ACCESS = "unauthorized_access"
DATA_EXPORT = "data_export"
CONFIG_CHANGED = "config_changed"
# Management Token 相关
MANAGEMENT_TOKEN_CREATED = "management_token_created"
MANAGEMENT_TOKEN_UPDATED = "management_token_updated"
MANAGEMENT_TOKEN_DELETED = "management_token_deleted"
MANAGEMENT_TOKEN_USED = "management_token_used"
MANAGEMENT_TOKEN_EXPIRED = "management_token_expired"
MANAGEMENT_TOKEN_IP_BLOCKED = "management_token_ip_blocked"
class AuditLog(Base):
"""审计日志模型"""
__tablename__ = "audit_logs"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
event_type = Column(String(50), nullable=False, index=True)
user_id = Column(
String(36), ForeignKey("users.id", ondelete="SET NULL"), nullable=True, index=True
)
api_key_id = Column(String(36), nullable=True)
# 事件详情
description = Column(Text, nullable=False)
ip_address = Column(String(45), nullable=True)
user_agent = Column(String(500), nullable=True)
request_id = Column(String(100), nullable=True, index=True)
# 相关数据
event_metadata = Column(JSON, nullable=True)
# 响应信息
status_code = Column(Integer, nullable=True)
error_message = Column(Text, nullable=True)
# 时间戳
created_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
nullable=False,
index=True,
)
# 关系
user = relationship("User", back_populates="audit_logs")
class GeminiFileMapping(Base):
"""
Gemini Files API 文件与 Provider Key 的映射关系
用于持久化存储 file_id -> key_id 的绑定关系,
确保后续 generateContent 请求使用上传时的同一 Key。
Gemini 文件有 48 小时有效期,此表中的记录也会在过期后被清理。
"""
__tablename__ = "gemini_file_mappings"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
# 文件名(如 files/abc123xyz
file_name = Column(String(255), nullable=False, unique=True, index=True)
# Provider Key ID关联到 provider_api_keys 表)
key_id = Column(
String(36),
ForeignKey("provider_api_keys.id", ondelete="CASCADE"),
nullable=False,
index=True,
)
# 用户 ID用于权限验证可选
user_id = Column(
String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=True, index=True
)
# 文件元数据(可选,用于调试)
display_name = Column(String(255), nullable=True)
mime_type = Column(String(100), nullable=True)
# 源文件哈希(用于关联相同源文件的不同上传,可选)
# 当同一源文件上传到多个 Key 时,可通过此字段找到所有等效文件
source_hash = Column(String(64), nullable=True, index=True)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
# 过期时间Gemini 文件 48 小时后过期)
expires_at = Column(DateTime(timezone=True), nullable=False, index=True)
# 关系
key = relationship("ProviderAPIKey")
user = relationship("User")
__table_args__ = (
Index("idx_gemini_file_mappings_expires", "expires_at"),
Index("idx_gemini_file_mappings_source_hash", "source_hash"),
)

516
src/models/model.py Normal file
View File

@@ -0,0 +1,516 @@
"""
模型相关数据库模型
包含: GlobalModel, Model, BillingRule, DimensionCollector
"""
from __future__ import annotations
import uuid
from datetime import datetime, timezone
from sqlalchemy import (
JSON,
Boolean,
CheckConstraint,
Column,
DateTime,
Float,
ForeignKey,
Index,
Integer,
String,
Text,
UniqueConstraint,
text,
)
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import relationship
from ._base import Base
class GlobalModel(Base):
"""全局统一模型定义 - 包含价格和能力配置
设计原则:
- 定义模型的基本信息和价格配置(价格为必填项)
- Provider 级别的 Model 可以覆盖这些默认值
- 如果 Model 的价格/能力字段为空,则使用 GlobalModel 的值
"""
__tablename__ = "global_models"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
name = Column(String(100), unique=True, nullable=False, index=True) # 统一模型名(唯一)
display_name = Column(String(100), nullable=False)
# 按次计费配置(每次请求的固定费用,美元)- 可选,与按 token 计费叠加
default_price_per_request = Column(Float, nullable=True, default=None) # 每次请求固定费用
# 统一阶梯计费配置JSON格式- 必填
# 固定价格也用单阶梯表示: {"tiers": [{"up_to": null, "input_price_per_1m": X, ...}]}
# 结构示例:
# {
# "tiers": [
# {
# "up_to": 128000, # 阶梯上限tokensnull 表示无上限
# "input_price_per_1m": 2.50,
# "output_price_per_1m": 10.00,
# "cache_creation_price_per_1m": 3.75, # 可选
# "cache_read_price_per_1m": 0.30, # 可选
# "cache_ttl_pricing": [ # 可选:按缓存时长分价格
# {"ttl_minutes": 5, "cache_read_price_per_1m": 0.30},
# {"ttl_minutes": 60, "cache_read_price_per_1m": 0.50}
# ]
# },
# {"up_to": null, "input_price_per_1m": 1.25, ...}
# ]
# }
default_tiered_pricing = Column(JSON, nullable=False)
# Key 能力配置 - 模型支持的能力列表(如 ["cache_1h", "context_1m"]
# Key 只能启用模型支持的能力
supported_capabilities = Column(JSON, nullable=True, default=list)
# 模型配置JSON格式- 包含能力、规格、元信息等
# 结构示例:
# {
# # 能力配置
# "streaming": true,
# "vision": true,
# "function_calling": true,
# "extended_thinking": false,
# "image_generation": false,
# # 规格参数
# "context_limit": 200000,
# "output_limit": 8192,
# # 元信息
# "description": "...",
# "icon_url": "...",
# "official_url": "...",
# "knowledge_cutoff": "2024-04",
# "family": "claude-3.5",
# "release_date": "2024-10-22",
# "input_modalities": ["text", "image"],
# "output_modalities": ["text"],
# }
config = Column(JSONB, nullable=True, default=dict)
# 状态
is_active = Column(Boolean, default=True, nullable=False)
# 统计计数器(优化性能,避免实时查询)
usage_count = Column(Integer, default=0, nullable=False, index=True)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
models = relationship("Model", back_populates="global_model")
class Model(Base):
"""Provider 模型配置表 - Provider 如何使用某个 GlobalModel
设计原则:
- Model 表示 Provider 对某个模型的具体实现
- global_model_id 可为空:
- 为空时:模型尚未关联到 GlobalModel不参与路由
- 不为空时:模型已关联 GlobalModel参与路由
- provider_model_name 是 Provider 侧的实际模型名称 (可能与 GlobalModel.name 不同)
- 价格和能力配置可为空,为空时使用 GlobalModel 的默认值
"""
__tablename__ = "models"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
provider_id = Column(String(36), ForeignKey("providers.id"), nullable=False)
# 可为空NULL 表示未关联,不参与路由;非 NULL 表示已关联,参与路由
global_model_id = Column(String(36), ForeignKey("global_models.id"), nullable=True, index=True)
# Provider 映射配置
provider_model_name = Column(String(200), nullable=False) # Provider 侧的主模型名称
# 模型名称映射列表(带优先级),用于同一模型在 Provider 侧有多个名称变体的场景
# 格式: [{"name": "Claude-Sonnet-4.5", "priority": 1}, {"name": "Claude-Sonnet-4-5", "priority": 2}]
# 为空时只使用 provider_model_name
provider_model_mappings = Column(JSON, nullable=True, default=None)
# 按次计费配置(每次请求的固定费用,美元)- 可为空,为空时使用 GlobalModel 的默认值
price_per_request = Column(Float, nullable=True) # 每次请求固定费用
# 阶梯计费配置JSON格式- 可为空,为空时使用 GlobalModel 的默认值
tiered_pricing = Column(JSON, nullable=True, default=None)
# Provider 能力配置 - 可为空,为空时使用 GlobalModel 的默认值
supports_vision = Column(Boolean, nullable=True)
supports_function_calling = Column(Boolean, nullable=True)
supports_streaming = Column(Boolean, nullable=True)
supports_extended_thinking = Column(Boolean, nullable=True)
supports_image_generation = Column(Boolean, nullable=True)
# 状态
is_active = Column(Boolean, default=True, nullable=False)
is_available = Column(Boolean, default=True) # 是否当前可用
# 扩展配置
config = Column(JSON, nullable=True)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
provider = relationship("Provider", back_populates="models")
global_model = relationship("GlobalModel", back_populates="models")
# 唯一约束:同一个提供商下的 provider_model_name 不能重复
__table_args__ = (
UniqueConstraint("provider_id", "provider_model_name", name="uq_provider_model"),
)
# 辅助方法:获取有效的阶梯计费配置
def get_effective_tiered_pricing(self) -> dict | None:
"""获取有效的阶梯计费配置"""
if self.tiered_pricing is not None:
return self.tiered_pricing
if self.global_model:
return self.global_model.default_tiered_pricing
return None
def _get_first_tier(self) -> dict | None:
"""获取第一个阶梯(用于获取默认价格)"""
tiered = self.get_effective_tiered_pricing()
if tiered and tiered.get("tiers"):
return tiered["tiers"][0]
return None
def get_effective_input_price(self) -> float:
"""获取有效的输入价格(从第一个阶梯)"""
tier = self._get_first_tier()
if tier:
return tier.get("input_price_per_1m", 0.0)
return 0.0
def get_effective_output_price(self) -> float:
"""获取有效的输出价格(从第一个阶梯)"""
tier = self._get_first_tier()
if tier:
return tier.get("output_price_per_1m", 0.0)
return 0.0
def get_effective_cache_creation_price(self) -> float | None:
"""获取有效的缓存创建价格(从第一个阶梯)"""
tier = self._get_first_tier()
if tier:
return tier.get("cache_creation_price_per_1m")
return None
def get_effective_cache_read_price(self) -> float | None:
"""获取有效的缓存读取价格(从第一个阶梯)"""
tier = self._get_first_tier()
if tier:
return tier.get("cache_read_price_per_1m")
return None
def get_effective_1h_cache_creation_price(self) -> float | None:
"""获取有效的 1h 缓存创建价格(从第一个阶梯)"""
tier = self._get_first_tier()
if tier:
cache_ttl_pricing = tier.get("cache_ttl_pricing") or []
for ttl_entry in cache_ttl_pricing:
if ttl_entry.get("ttl_minutes") == 60:
return ttl_entry.get("cache_creation_price_per_1m")
return None
def get_effective_price_per_request(self) -> float | None:
"""获取有效的按次计费价格"""
if self.price_per_request is not None:
return self.price_per_request
if self.global_model:
return self.global_model.default_price_per_request
return None
def _get_effective_capability(self, attr_name: str, default: bool = False) -> bool:
"""获取有效的能力配置(通用辅助方法)"""
local_value = getattr(self, attr_name, None)
if local_value is not None:
return bool(local_value)
if self.global_model:
config_key_map = {
"supports_vision": "vision",
"supports_function_calling": "function_calling",
"supports_streaming": "streaming",
"supports_extended_thinking": "extended_thinking",
"supports_image_generation": "image_generation",
}
config_key = config_key_map.get(attr_name)
if config_key:
global_config = getattr(self.global_model, "config", None)
if isinstance(global_config, dict):
global_value = global_config.get(config_key)
if global_value is not None:
return bool(global_value)
return default
def get_effective_supports_vision(self) -> bool:
return self._get_effective_capability("supports_vision", False)
def get_effective_supports_function_calling(self) -> bool:
return self._get_effective_capability("supports_function_calling", False)
def get_effective_supports_streaming(self) -> bool:
return self._get_effective_capability("supports_streaming", True)
def get_effective_supports_extended_thinking(self) -> bool:
return self._get_effective_capability("supports_extended_thinking", False)
def get_effective_supports_image_generation(self) -> bool:
return self._get_effective_capability("supports_image_generation", False)
def get_effective_config(self) -> dict | None:
"""获取有效的 config合并 Model 和 GlobalModel 的 config
合并策略:
- GlobalModel.config 作为基础
- Model.config 覆盖 GlobalModel.config
- 深度合并 billing 子字段
"""
global_config = {}
if self.global_model and self.global_model.config:
global_config = dict(self.global_model.config)
if not self.config:
return global_config if global_config else None
# 深度合并 config
result = dict(global_config)
for key, value in self.config.items():
if key == "billing" and isinstance(value, dict) and isinstance(result.get(key), dict):
# 深度合并 billing
result[key] = {**result[key], **value}
else:
result[key] = value
return result if result else None
def select_provider_model_name(
self, affinity_key: str | None = None, api_format: str | None = None
) -> str:
"""按优先级选择要使用的 Provider 模型名称
如果配置了 provider_model_mappings按优先级选择数字越小越优先
相同优先级的映射通过哈希分散实现负载均衡(与 Key 调度策略一致);
否则返回 provider_model_name。
Args:
affinity_key: 用于哈希分散的亲和键(如用户 API Key 哈希),确保同一用户稳定选择同一映射
api_format: 当前请求的 endpoint signature"openai:chat"),用于过滤适用的映射
"""
import hashlib
if not self.provider_model_mappings:
return self.provider_model_name
raw_mappings = self.provider_model_mappings
if not isinstance(raw_mappings, list) or len(raw_mappings) == 0:
return self.provider_model_name
mappings: list[dict] = []
for raw in raw_mappings:
if not isinstance(raw, dict):
continue
name = raw.get("name")
if not isinstance(name, str) or not name.strip():
continue
# 检查 api_formats 作用域(如果配置了且当前有 api_format
mapping_api_formats = raw.get("api_formats")
if api_format and mapping_api_formats:
# 如果配置了作用域,只有匹配时才生效
if isinstance(mapping_api_formats, list):
target = str(api_format).strip().lower()
allowed = {str(fmt).strip().lower() for fmt in mapping_api_formats if fmt}
if target not in allowed:
continue
raw_priority = raw.get("priority", 1)
try:
priority = int(raw_priority)
except Exception:
priority = 1
if priority < 1:
priority = 1
mappings.append({"name": name.strip(), "priority": priority})
if not mappings:
return self.provider_model_name
# 按优先级排序(数字越小越优先)
sorted_mappings = sorted(mappings, key=lambda x: x["priority"])
# 获取最高优先级(最小数字)
highest_priority = sorted_mappings[0]["priority"]
# 获取所有最高优先级的映射
top_priority_mappings = [
mapping for mapping in sorted_mappings if mapping["priority"] == highest_priority
]
# 如果有多个相同优先级的映射,通过哈希分散选择
if len(top_priority_mappings) > 1 and affinity_key:
# 为每个映射计算哈希得分,选择得分最小的
def hash_score(mapping: dict) -> int:
combined = f"{affinity_key}:{mapping['name']}"
return int(hashlib.md5(combined.encode()).hexdigest(), 16)
selected = min(top_priority_mappings, key=hash_score)
elif len(top_priority_mappings) > 1:
# 没有 affinity_key 时,使用确定性选择(按名称排序后取第一个)
# 避免随机选择导致同一请求重试时选择不同的模型名称
selected = min(top_priority_mappings, key=lambda x: x["name"])
else:
selected = top_priority_mappings[0]
return selected["name"]
def get_all_provider_model_names(self) -> list[str]:
"""获取所有可用的 Provider 模型名称(主名称 + 映射名称)"""
names = [self.provider_model_name]
if self.provider_model_mappings:
for mapping in self.provider_model_mappings:
if isinstance(mapping, dict) and mapping.get("name"):
names.append(mapping["name"])
return names
class BillingRule(Base):
"""计费规则表(单条 formula 规则,支持 Model 覆盖 GlobalModel"""
__tablename__ = "billing_rules"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# 规则关联(两者必有其一)
global_model_id = Column(
String(36), ForeignKey("global_models.id", ondelete="CASCADE"), nullable=True, index=True
)
model_id = Column(
String(36), ForeignKey("models.id", ondelete="CASCADE"), nullable=True, index=True
)
name = Column(String(100), nullable=False)
# 注CLI 在计费域里恒等于 chat不单独存 "cli"
task_type = Column(String(20), nullable=False, default="chat")
# Formula 表达式及其配置
expression = Column(Text, nullable=False)
variables = Column(JSONB, nullable=False, default=dict)
dimension_mappings = Column(JSONB, nullable=False, default=dict)
is_enabled = Column(Boolean, nullable=False, default=True)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
global_model = relationship("GlobalModel", foreign_keys=[global_model_id])
model = relationship("Model", foreign_keys=[model_id])
__table_args__ = (
CheckConstraint(
"(global_model_id IS NOT NULL AND model_id IS NULL) OR "
"(global_model_id IS NULL AND model_id IS NOT NULL)",
name="chk_billing_rules_model_ref",
),
# 同级同 task_type 只允许一条启用规则partial unique index
Index(
"uq_billing_rules_global_model_task",
"global_model_id",
"task_type",
unique=True,
postgresql_where=text("is_enabled = TRUE AND global_model_id IS NOT NULL"),
),
Index(
"uq_billing_rules_model_task",
"model_id",
"task_type",
unique=True,
postgresql_where=text("is_enabled = TRUE AND model_id IS NOT NULL"),
),
)
class DimensionCollector(Base):
"""维度收集器配置表(从请求/响应/元数据/派生计算收集维度)。"""
__tablename__ = "dimension_collectors"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
api_format = Column(String(50), nullable=False)
task_type = Column(String(20), nullable=False)
dimension_name = Column(String(100), nullable=False)
# 来源配置
# - response / request / metadata / computed
source_type = Column(String(20), nullable=False)
source_path = Column(String(200), nullable=True) # computed 允许为空
# 值类型与转换
value_type = Column(String(20), nullable=False, default="float") # float/int/string
transform_expression = Column(Text, nullable=True) # computed 时为派生公式
default_value = Column(String(100), nullable=True)
priority = Column(Integer, nullable=False, default=0)
is_enabled = Column(Boolean, nullable=False, default=True)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
__table_args__ = (
CheckConstraint(
"(source_type = 'computed' AND source_path IS NULL AND transform_expression IS NOT NULL) OR "
"(source_type != 'computed' AND source_path IS NOT NULL)",
name="chk_dimension_collectors_source_config",
),
# 同维度 + 同优先级 + enabled 才唯一(允许禁用旧配置后重建)
Index(
"uq_dimension_collectors_enabled",
"api_format",
"task_type",
"dimension_name",
"priority",
unique=True,
postgresql_where=text("is_enabled = TRUE"),
),
)

426
src/models/provider.py Normal file
View File

@@ -0,0 +1,426 @@
"""
提供商相关数据库模型
包含: Provider, ProviderEndpoint, ProxyNodeStatus, ProxyNode, ProviderAPIKey
"""
from __future__ import annotations
import uuid
from datetime import datetime, timezone
from enum import Enum as PyEnum
from sqlalchemy import (
JSON,
BigInteger,
Boolean,
Column,
DateTime,
Enum,
Float,
ForeignKey,
Index,
Integer,
String,
Text,
UniqueConstraint,
)
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import relationship
from src.core.enums import ProviderBillingType
from ._base import Base
class Provider(Base):
"""提供商配置表"""
__tablename__ = "providers"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
name = Column(String(100), unique=True, nullable=False, index=True) # 提供商名称(唯一)
description = Column(Text, nullable=True) # 提供商描述
website = Column(String(500), nullable=True) # 主站网站
# Provider 类型(用于模板化固定 Provider / 自定义 Provider
# - custom: 自定义
# - claude_code / codex / gemini_cli / antigravity: 固定类型
provider_type = Column(String(20), default="custom", nullable=False)
# 计费类型配置
billing_type = Column(
Enum(
ProviderBillingType,
name="providerbillingtype",
create_type=False,
values_callable=lambda x: [e.value for e in x],
),
default=ProviderBillingType.PAY_AS_YOU_GO,
nullable=False,
)
# 月卡配置
monthly_quota_usd = Column(Float, nullable=True) # 月卡总额度
monthly_used_usd = Column(Float, default=0.0) # 本月已用额度
quota_reset_day = Column(Integer, default=30) # 额度重置周期(天数)例如7=每周30=每月
quota_last_reset_at = Column(DateTime(timezone=True), nullable=True) # 上次额度重置时间
quota_expires_at = Column(DateTime(timezone=True), nullable=True) # 月卡过期时间
# 提供商优先级 (数字越小越优先,用于提供商优先模式下的 Provider 排序)
# 0-10: 急需消耗(如即将过期的月卡)
# 11-50: 优先消耗(月卡)
# 51-100: 正常消费(按量付费)
# 101+: 备用(高成本或限制严格的)
provider_priority = Column(Integer, default=100)
# 格式转换时是否保持优先级(默认 False
# - False: 需要格式转换时,该提供商的候选会被降级到不需要转换的候选之后
# - True: 即使需要格式转换,也保持原优先级排名
# 注意:如果系统配置 keep_priority_on_conversion=true此字段被忽略所有提供商都保持优先级
keep_priority_on_conversion = Column(Boolean, default=False, nullable=False)
# 是否允许格式转换(默认 True
# - True: 该提供商可以作为格式转换的目标(如 OpenAI 客户端请求可以路由到此 Gemini 提供商)
# - False: 该提供商不接受需要格式转换的请求
# 优先级逻辑:
# - 全局开关 ON -> 强制允许所有提供商的格式转换(忽略此字段)
# - 全局开关 OFF -> 由此字段决定是否允许该提供商的格式转换
enable_format_conversion = Column(Boolean, default=False, nullable=False)
# 状态
is_active = Column(Boolean, default=True, nullable=False)
# 限制
concurrent_limit = Column(Integer, nullable=True) # 并发请求限制
# 请求配置
max_retries = Column(Integer, default=2, nullable=True) # 最大重试次数
proxy = Column(JSONB, nullable=True) # 代理配置: {url, username, password, enabled}
# 超时配置(秒),为 None 时使用全局配置
stream_first_byte_timeout = Column(Float, nullable=True) # 流式请求首字节超时
request_timeout = Column(Float, nullable=True) # 非流式请求整体超时
# 配置
config = Column(JSON, nullable=True) # 额外配置如Azure deployment name等
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
models = relationship("Model", back_populates="provider", cascade="all, delete-orphan")
endpoints = relationship(
"ProviderEndpoint", back_populates="provider", cascade="all, delete-orphan"
)
api_keys = relationship(
"ProviderAPIKey", back_populates="provider", cascade="all, delete-orphan"
)
api_key_mappings = relationship(
"ApiKeyProviderMapping", back_populates="provider", cascade="all, delete-orphan"
)
usage_tracking = relationship(
"ProviderUsageTracking", back_populates="provider", cascade="all, delete-orphan"
)
class ProviderEndpoint(Base):
"""提供商端点 - 一个提供商可以有多个 API 格式端点"""
__tablename__ = "provider_endpoints"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
provider_id = Column(String(36), ForeignKey("providers.id", ondelete="CASCADE"), nullable=False)
# API 格式和配置
# 新模式:存储 endpoint signature keyfamily:kind如 "openai:chat"
api_format = Column(String(50), nullable=False)
# 新架构字段Phase 1/3用于将 api_format 拆分为结构化维度
api_family = Column(String(50), nullable=True) # openai/claude/gemini
endpoint_kind = Column(String(50), nullable=True) # chat/cli/video/...
base_url = Column(String(500), nullable=False)
# 请求配置
header_rules = Column(JSON, nullable=True) # 请求头规则 [{action, key, value, from, to}]
body_rules = Column(JSON, nullable=True) # 请求体规则 [{action, path, value, from, to}]
max_retries = Column(Integer, default=2) # 最大重试次数
# 状态
is_active = Column(Boolean, default=True, nullable=False)
# 路径配置
custom_path = Column(
String(200), nullable=True
) # 自定义请求路径,为空则使用 API 格式的默认路径
# 额外配置
config = Column(JSON, nullable=True) # 端点特定配置(不推荐使用,优先使用专用字段)
# 格式转换配置
format_acceptance_config = Column(
JSON,
nullable=True,
default=None,
comment="格式接受策略配置(跨格式转换开关/白黑名单等)",
)
# 代理配置
proxy = Column(JSONB, nullable=True) # 代理配置: {url, username, password}
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
provider = relationship("Provider", back_populates="endpoints")
# 唯一约束和索引在表定义后
__table_args__ = (
UniqueConstraint("provider_id", "api_format", name="uq_provider_api_format"),
Index("idx_endpoint_format_active", "api_format", "is_active"),
Index("idx_provider_family_kind", "provider_id", "api_family", "endpoint_kind"),
)
class ProxyNodeStatus(PyEnum):
"""代理节点状态"""
ONLINE = "online"
UNHEALTHY = "unhealthy"
OFFLINE = "offline"
class ProxyNode(Base):
"""代理节点表aether-proxy 自动注册 + 手动添加)"""
__tablename__ = "proxy_nodes"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
name = Column(String(100), nullable=False) # 节点名
ip = Column(String(512), nullable=False) # 公网 IP 或手动节点的主机名(含协议前缀)
port = Column(Integer, nullable=False) # 代理端口
region = Column(String(100), nullable=True) # 区域标签
# 手动节点专用字段
is_manual = Column(Boolean, default=False, nullable=False, comment="是否为手动添加的代理节点")
proxy_url = Column(String(500), nullable=True, comment="手动节点的完整代理 URL")
proxy_username = Column(String(255), nullable=True, comment="手动节点的代理用户名")
proxy_password = Column(String(500), nullable=True, comment="手动节点的代理密码")
status = Column(
Enum(
ProxyNodeStatus,
name="proxynodestatus",
create_type=False,
values_callable=lambda x: [e.value for e in x],
),
default=ProxyNodeStatus.ONLINE,
nullable=False,
)
registered_by = Column(
String(36),
ForeignKey("users.id", ondelete="SET NULL"),
nullable=True,
comment="注册该节点的管理员用户 ID可空",
)
last_heartbeat_at = Column(DateTime(timezone=True), nullable=True)
heartbeat_interval = Column(Integer, default=30, nullable=False)
# 性能指标(心跳上报)
active_connections = Column(Integer, default=0, nullable=False)
total_requests = Column(BigInteger, default=0, nullable=False)
avg_latency_ms = Column(Float, nullable=True)
# TLS 加密
tls_enabled = Column(Boolean, default=False, nullable=False, comment="是否启用 TLS 加密")
tls_cert_fingerprint = Column(
String(128), nullable=True, comment="TLS 证书 SHA-256 指纹hex"
)
# 硬件信息注册时上报JSON 可扩展)
hardware_info = Column(
JSON,
nullable=True,
comment="硬件信息 (cpu_cores, total_memory_mb, os_info, fd_limit, ...)",
)
estimated_max_concurrency = Column(
Integer, nullable=True, comment="基于硬件估算的最大并发连接数"
)
# 管理端远程配置(通过心跳下发给 aether-proxy
remote_config = Column(
JSON,
nullable=True,
comment="管理端下发的远程配置 (allowed_ports, log_level, heartbeat_interval, timestamp_tolerance)",
)
config_version = Column(
Integer, default=0, nullable=False, comment="远程配置版本号,每次更新 +1"
)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
__table_args__ = (UniqueConstraint("ip", "port", name="uq_proxy_node_ip_port"),)
class ProviderAPIKey(Base):
"""Provider API密钥表 - 直接归属于 Provider支持多种 API 格式"""
__tablename__ = "provider_api_keys"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
# 外键关系 - 直接关联 Provider
provider_id = Column(
String(36), ForeignKey("providers.id", ondelete="CASCADE"), nullable=False, index=True
)
# API 格式支持列表(核心字段)
# None 表示支持所有格式(兼容历史数据),空列表 [] 表示不支持任何格式
api_formats = Column(JSON, nullable=True, default=list) # ["claude:chat", "claude:cli"]
# 认证类型
# - "api_key": 标准 API Key 认证(默认)
# - "vertex_ai": Google Vertex AI 认证Service Account JSON
# - 未来可扩展oauth2, azure_ad, aws_iam 等
auth_type = Column(String(20), default="api_key", nullable=False)
# API密钥加密存储
# - auth_type="api_key" 时:存储 API Key 字符串
# - auth_type="vertex_ai" 等:可为空,敏感凭证存在 auth_config 中
api_key = Column(Text, nullable=False) # 使用 Text 支持加密后的 OAuth token
# 认证配置(加密存储)
# - auth_type="api_key" 时:可为空
# - auth_type="vertex_ai" 时:存储加密后的 Service Account JSON
# - auth_type="oauth2" 时:存储加密后的 {client_id, client_secret, token_url, scope}
auth_config = Column(Text, nullable=True)
name = Column(String(100), nullable=False) # 密钥名称(必填,用于识别)
note = Column(String(500), nullable=True) # 备注说明(可选)
# 成本计算
rate_multipliers = Column(
JSON, nullable=True
) # 按 endpoint signature 的成本倍率 {"claude:cli": 1.0, "openai:cli": 0.8}
# 优先级配置 (数字越小越优先)
internal_priority = Column(
Integer, default=50
) # Endpoint 内部优先级(用于提供商优先模式,同 Endpoint 内 Keys 的排序,同优先级参与负载均衡)
global_priority_by_format = Column(
JSON, nullable=True
) # 按 endpoint signature 的全局优先级 {"claude:chat": 1, "claude:cli": 2}
# RPM 限制配置(自适应学习)
# rpm_limit 决定 RPM 控制模式:
# - NULL: 自适应模式,系统自动学习并调整(使用 learned_rpm_limit
# - 数字: 固定限制模式,使用用户指定的值
rpm_limit = Column(Integer, nullable=True, default=None)
# 模型权限控制
allowed_models = Column(JSON, nullable=True) # 允许使用的模型列表null = 支持所有模型)
# Key 能力标签
capabilities = Column(JSON, nullable=True) # Key 拥有的能力
# 示例: {"cache_1h": true, "context_1m": true}
# 自适应 RPM 调整(仅当 rpm_limit = NULL 时生效)
learned_rpm_limit = Column(Integer, nullable=True) # 学习到的 RPM 限制(自适应模式下的有效值)
concurrent_429_count = Column(Integer, default=0, nullable=False) # 因并发导致的429次数
rpm_429_count = Column(Integer, default=0, nullable=False) # 因RPM导致的429次数
last_429_at = Column(DateTime(timezone=True), nullable=True) # 最后429时间
last_429_type = Column(String(50), nullable=True) # 最后429类型: concurrent/rpm/unknown
last_rpm_peak = Column(Integer, nullable=True) # 触发429时的RPM峰值
adjustment_history = Column(JSON, nullable=True) # RPM调整历史
# 基于滑动窗口的利用率追踪
utilization_samples = Column(
JSON, nullable=True
) # 利用率采样窗口 [{"ts": timestamp, "util": 0.8}, ...]
last_probe_increase_at = Column(DateTime(timezone=True), nullable=True) # 上次探测性扩容时间
# 健康度追踪(按 endpoint signature 存储)
# 结构: {"claude:chat": {"health_score": 1.0, "consecutive_failures": 0, ...}, ...}
health_by_format = Column(JSON, nullable=True, default=dict)
# 缓存与熔断配置
cache_ttl_minutes = Column(
Integer, default=5, nullable=False
) # 缓存TTL(分钟)0表示不支持缓存默认5分钟
max_probe_interval_minutes = Column(
Integer, default=32, nullable=False
) # 最大探测间隔(分钟)默认32分钟硬上限
# 熔断器状态(按 endpoint signature 存储)
# 结构: {"claude:chat": {"open": false, "open_at": null, ...}, ...}
circuit_breaker_by_format = Column(JSON, nullable=True, default=dict)
# 使用统计
request_count = Column(Integer, default=0) # 请求次数
success_count = Column(Integer, default=0) # 成功次数
error_count = Column(Integer, default=0) # 错误次数
total_response_time_ms = Column(Integer, default=0) # 总响应时间(用于计算平均值)
last_used_at = Column(DateTime(timezone=True), nullable=True) # 最后使用时间
last_error_at = Column(DateTime(timezone=True), nullable=True) # 最后错误时间
last_error_msg = Column(Text, nullable=True) # 最后错误信息
# 状态
is_active = Column(Boolean, default=True, nullable=False)
expires_at = Column(DateTime(timezone=True), nullable=True) # 过期时间
# 自动获取模型配置
auto_fetch_models = Column(Boolean, default=False, nullable=False) # 是否启用自动获取模型
last_models_fetch_at = Column(DateTime(timezone=True), nullable=True) # 最后获取时间
last_models_fetch_error = Column(Text, nullable=True) # 最后获取错误信息
locked_models = Column(JSON, nullable=True) # 被锁定的模型列表(刷新时不会被删除)
# 模型过滤规则(支持 * 和 ? 通配符,如 "gpt-*", "claude-?-sonnet"
model_include_patterns = Column(JSON, nullable=True) # 包含规则列表,空表示不过滤(包含所有)
model_exclude_patterns = Column(JSON, nullable=True) # 排除规则列表,空表示不排除
# 上游元数据(由响应头解析器采集,如 Codex 额度信息)
upstream_metadata = Column(JSON, nullable=True, default=dict)
# OAuth 失效状态(账号被封、授权撤销、刷新失败等)
oauth_invalid_at = Column(DateTime(timezone=True), nullable=True) # 失效时间
oauth_invalid_reason = Column(String(255), nullable=True) # 失效原因
# Key 级别的代理配置(覆盖 Provider 级别的代理设置)
# 结构: {"node_id": "xxx", "enabled": true} 或 {"url": "socks5://...", "enabled": true}
# null 表示使用 Provider 级别代理(默认行为)
proxy = Column(JSON, nullable=True, default=None)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
provider = relationship("Provider", back_populates="api_keys")

461
src/models/stats.py Normal file
View File

@@ -0,0 +1,461 @@
"""
统计数据相关数据库模型
包含: StatsBaseMixin, StatsHourly, StatsHourlyUser, StatsHourlyModel, StatsHourlyProvider,
StatsDaily, StatsDailyModel, StatsDailyProvider, StatsDailyApiKey, StatsDailyError,
StatsSummary, StatsUserDaily
"""
from __future__ import annotations
import uuid
from datetime import datetime, timezone
from sqlalchemy import (
BigInteger,
Boolean,
Column,
DateTime,
Float,
ForeignKey,
Index,
Integer,
String,
UniqueConstraint,
)
from sqlalchemy.orm import relationship
from ._base import Base
class StatsBaseMixin:
"""统计表公共字段 Mixin"""
total_requests = Column(Integer, default=0, nullable=False)
input_tokens = Column(BigInteger, default=0, nullable=False)
output_tokens = Column(BigInteger, default=0, nullable=False)
total_cost = Column(Float, default=0.0, nullable=False)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
class StatsHourly(Base):
"""小时级统计快照 - 用于时间序列查询"""
__tablename__ = "stats_hourly"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# 小时起点 (UTC)
hour_utc = Column(DateTime(timezone=True), nullable=False, unique=True, index=True)
# 请求统计
total_requests = Column(Integer, default=0, nullable=False)
success_requests = Column(Integer, default=0, nullable=False)
error_requests = Column(Integer, default=0, nullable=False)
# Token 统计
input_tokens = Column(BigInteger, default=0, nullable=False)
output_tokens = Column(BigInteger, default=0, nullable=False)
cache_creation_tokens = Column(BigInteger, default=0, nullable=False)
cache_read_tokens = Column(BigInteger, default=0, nullable=False)
# 成本统计 (USD)
total_cost = Column(Float, default=0.0, nullable=False)
actual_total_cost = Column(Float, default=0.0, nullable=False)
# 性能统计
avg_response_time_ms = Column(Float, default=0.0, nullable=False)
# 完成标记
is_complete = Column(Boolean, default=False, nullable=False)
aggregated_at = Column(DateTime(timezone=True), nullable=True)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
__table_args__ = (Index("idx_stats_hourly_hour", "hour_utc"),)
class StatsHourlyUser(StatsBaseMixin, Base):
"""小时级用户维度统计"""
__tablename__ = "stats_hourly_user"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
hour_utc = Column(DateTime(timezone=True), nullable=False, index=True)
user_id = Column(String(36), nullable=False, index=True)
success_requests = Column(Integer, default=0, nullable=False)
error_requests = Column(Integer, default=0, nullable=False)
__table_args__ = (
UniqueConstraint("hour_utc", "user_id", name="uq_stats_hourly_user"),
Index("idx_stats_hourly_user_hour", "hour_utc"),
Index("idx_stats_hourly_user_user_hour", "user_id", "hour_utc"),
)
class StatsHourlyModel(StatsBaseMixin, Base):
"""小时级模型维度统计"""
__tablename__ = "stats_hourly_model"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
hour_utc = Column(DateTime(timezone=True), nullable=False, index=True)
model = Column(String(100), nullable=False, index=True)
avg_response_time_ms = Column(Float, default=0.0, nullable=False)
__table_args__ = (
UniqueConstraint("hour_utc", "model", name="uq_stats_hourly_model"),
Index("idx_stats_hourly_model_hour", "hour_utc"),
Index("idx_stats_hourly_model_model_hour", "model", "hour_utc"),
)
class StatsHourlyProvider(StatsBaseMixin, Base):
"""小时级提供商维度统计"""
__tablename__ = "stats_hourly_provider"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
hour_utc = Column(DateTime(timezone=True), nullable=False, index=True)
provider_name = Column(String(100), nullable=False, index=True)
__table_args__ = (
UniqueConstraint("hour_utc", "provider_name", name="uq_stats_hourly_provider"),
Index("idx_stats_hourly_provider_hour", "hour_utc"),
)
class StatsDaily(Base):
"""每日统计快照 - 用于快速查询历史数据"""
__tablename__ = "stats_daily"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# 统计日期 (UTC)
date = Column(DateTime(timezone=True), nullable=False, unique=True, index=True)
# 请求统计
total_requests = Column(Integer, default=0, nullable=False)
success_requests = Column(Integer, default=0, nullable=False)
error_requests = Column(Integer, default=0, nullable=False)
# Token 统计
input_tokens = Column(BigInteger, default=0, nullable=False)
output_tokens = Column(BigInteger, default=0, nullable=False)
cache_creation_tokens = Column(BigInteger, default=0, nullable=False)
cache_read_tokens = Column(BigInteger, default=0, nullable=False)
# 成本统计 (USD)
total_cost = Column(Float, default=0.0, nullable=False)
actual_total_cost = Column(Float, default=0.0, nullable=False) # 倍率后成本
input_cost = Column(Float, default=0.0, nullable=False)
output_cost = Column(Float, default=0.0, nullable=False)
cache_creation_cost = Column(Float, default=0.0, nullable=False)
cache_read_cost = Column(Float, default=0.0, nullable=False)
# 性能统计
avg_response_time_ms = Column(Float, default=0.0, nullable=False)
p50_response_time_ms = Column(Integer, nullable=True)
p90_response_time_ms = Column(Integer, nullable=True)
p99_response_time_ms = Column(Integer, nullable=True)
p50_first_byte_time_ms = Column(Integer, nullable=True)
p90_first_byte_time_ms = Column(Integer, nullable=True)
p99_first_byte_time_ms = Column(Integer, nullable=True)
fallback_count = Column(Integer, default=0, nullable=False) # Provider 切换次数
# 使用维度统计
unique_models = Column(Integer, default=0, server_default="0", nullable=False)
unique_providers = Column(Integer, default=0, server_default="0", nullable=False)
# 完成标记
is_complete = Column(Boolean, default=False, nullable=False)
aggregated_at = Column(DateTime(timezone=True), nullable=True)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
class StatsDailyModel(Base):
"""每日模型统计快照 - 用于快速查询每日模型维度数据"""
__tablename__ = "stats_daily_model"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# 统计日期 (UTC)
date = Column(DateTime(timezone=True), nullable=False, index=True)
# 模型名称
model = Column(String(100), nullable=False)
# 请求统计
total_requests = Column(Integer, default=0, nullable=False)
# Token 统计
input_tokens = Column(BigInteger, default=0, nullable=False)
output_tokens = Column(BigInteger, default=0, nullable=False)
cache_creation_tokens = Column(BigInteger, default=0, nullable=False)
cache_read_tokens = Column(BigInteger, default=0, nullable=False)
# 成本统计 (USD)
total_cost = Column(Float, default=0.0, nullable=False)
# 性能统计
avg_response_time_ms = Column(Float, default=0.0, nullable=False)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 唯一约束:每个模型每天只有一条记录
__table_args__ = (
UniqueConstraint("date", "model", name="uq_stats_daily_model"),
Index("idx_stats_daily_model_date", "date"),
Index("idx_stats_daily_model_date_model", "date", "model"),
)
class StatsDailyProvider(Base):
"""每日供应商统计快照 - 用于快速查询每日供应商维度数据"""
__tablename__ = "stats_daily_provider"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# 统计日期 (UTC)
date = Column(DateTime(timezone=True), nullable=False, index=True)
# 供应商名称
provider_name = Column(String(100), nullable=False)
# 请求统计
total_requests = Column(Integer, default=0, nullable=False)
# Token 统计
input_tokens = Column(BigInteger, default=0, nullable=False)
output_tokens = Column(BigInteger, default=0, nullable=False)
cache_creation_tokens = Column(BigInteger, default=0, nullable=False)
cache_read_tokens = Column(BigInteger, default=0, nullable=False)
# 成本统计 (USD)
total_cost = Column(Float, default=0.0, nullable=False)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 唯一约束:每个供应商每天只有一条记录
__table_args__ = (
UniqueConstraint("date", "provider_name", name="uq_stats_daily_provider"),
Index("idx_stats_daily_provider_date", "date"),
Index("idx_stats_daily_provider_date_provider", "date", "provider_name"),
)
class StatsDailyApiKey(Base):
"""API Key 每日统计"""
__tablename__ = "stats_daily_api_key"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
api_key_id = Column(String(36), ForeignKey("api_keys.id", ondelete="CASCADE"), nullable=False)
date = Column(DateTime(timezone=True), nullable=False, index=True)
total_requests = Column(Integer, default=0, nullable=False)
success_requests = Column(Integer, default=0, nullable=False)
error_requests = Column(Integer, default=0, nullable=False)
input_tokens = Column(BigInteger, default=0, nullable=False)
output_tokens = Column(BigInteger, default=0, nullable=False)
cache_creation_tokens = Column(BigInteger, default=0, nullable=False)
cache_read_tokens = Column(BigInteger, default=0, nullable=False)
total_cost = Column(Float, default=0.0, nullable=False)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
__table_args__ = (
UniqueConstraint("api_key_id", "date", name="uq_stats_daily_api_key"),
Index("idx_stats_daily_api_key_date", "date"),
Index("idx_stats_daily_api_key_key_date", "api_key_id", "date"),
Index("idx_stats_daily_api_key_date_requests", "date", "total_requests"),
Index("idx_stats_daily_api_key_date_cost", "date", "total_cost"),
)
api_key = relationship("ApiKey")
class StatsDailyError(Base):
"""每日错误统计"""
__tablename__ = "stats_daily_error"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
date = Column(DateTime(timezone=True), nullable=False, index=True)
error_category = Column(String(50), nullable=False)
provider_name = Column(String(100), nullable=True)
model = Column(String(100), nullable=True)
count = Column(Integer, default=0, nullable=False)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
__table_args__ = (
UniqueConstraint(
"date",
"error_category",
"provider_name",
"model",
name="uq_stats_daily_error",
),
Index("idx_stats_daily_error_date", "date"),
Index("idx_stats_daily_error_category", "date", "error_category"),
)
class StatsSummary(Base):
"""全局统计汇总 - 单行记录,存储截止到昨天的累计数据"""
__tablename__ = "stats_summary"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# 统计截止日期 (不含当天)
cutoff_date = Column(DateTime(timezone=True), nullable=False)
# 累计请求统计
all_time_requests = Column(Integer, default=0, nullable=False)
all_time_success_requests = Column(Integer, default=0, nullable=False)
all_time_error_requests = Column(Integer, default=0, nullable=False)
# 累计 Token 统计
all_time_input_tokens = Column(BigInteger, default=0, nullable=False)
all_time_output_tokens = Column(BigInteger, default=0, nullable=False)
all_time_cache_creation_tokens = Column(BigInteger, default=0, nullable=False)
all_time_cache_read_tokens = Column(BigInteger, default=0, nullable=False)
# 累计成本统计 (USD)
all_time_cost = Column(Float, default=0.0, nullable=False)
all_time_actual_cost = Column(Float, default=0.0, nullable=False)
# 累计用户/API Key 统计 (快照)
total_users = Column(Integer, default=0, nullable=False)
active_users = Column(Integer, default=0, nullable=False)
total_api_keys = Column(Integer, default=0, nullable=False)
active_api_keys = Column(Integer, default=0, nullable=False)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
class StatsUserDaily(Base):
"""用户每日统计快照 - 用于用户仪表盘"""
__tablename__ = "stats_user_daily"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# 用户关联
user_id = Column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False)
# 统计日期 (UTC)
date = Column(DateTime(timezone=True), nullable=False, index=True)
# 请求统计
total_requests = Column(Integer, default=0, nullable=False)
success_requests = Column(Integer, default=0, nullable=False)
error_requests = Column(Integer, default=0, nullable=False)
# Token 统计
input_tokens = Column(BigInteger, default=0, nullable=False)
output_tokens = Column(BigInteger, default=0, nullable=False)
cache_creation_tokens = Column(BigInteger, default=0, nullable=False)
cache_read_tokens = Column(BigInteger, default=0, nullable=False)
# 成本统计 (USD)
total_cost = Column(Float, default=0.0, nullable=False)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 唯一约束:每个用户每天只有一条记录
__table_args__ = (
UniqueConstraint("user_id", "date", name="uq_stats_user_daily"),
Index("idx_stats_user_daily_user_date", "user_id", "date"),
)
# 关系
user = relationship("User")

243
src/models/usage.py Normal file
View File

@@ -0,0 +1,243 @@
"""
使用记录相关数据库模型
包含: Usage, RequestCandidate
"""
from __future__ import annotations
import uuid
from datetime import datetime, timezone
from typing import Any
from sqlalchemy import (
JSON,
Boolean,
Column,
DateTime,
Float,
ForeignKey,
Index,
Integer,
LargeBinary,
String,
Text,
UniqueConstraint,
)
from sqlalchemy.orm import relationship
from ._base import Base
class Usage(Base):
"""使用记录模型"""
__tablename__ = "usage"
__table_args__ = (
# Composite indexes for common query patterns (analytics / list pages)
Index("idx_usage_user_created", "user_id", "created_at"),
Index("idx_usage_apikey_created", "api_key_id", "created_at"),
Index("idx_usage_provider_model_created", "provider_name", "model", "created_at"),
Index("idx_usage_provider_created", "provider_name", "created_at"),
Index("idx_usage_model_created", "model", "created_at"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
user_id = Column(String(36), ForeignKey("users.id", ondelete="SET NULL"), nullable=True)
api_key_id = Column(String(36), ForeignKey("api_keys.id", ondelete="SET NULL"), nullable=True)
# 请求信息
request_id = Column(String(100), unique=True, index=True, nullable=False)
provider_name = Column(String(100), nullable=False) # Provider 名称(非外键)
model = Column(String(100), nullable=False)
target_model = Column(
String(100), nullable=True, comment="映射后的目标模型名(若无映射则为空)"
)
# Provider 侧追踪信息(记录最终成功的 Provider/Endpoint/Key
provider_id = Column(String(36), ForeignKey("providers.id", ondelete="SET NULL"), nullable=True)
provider_endpoint_id = Column(
String(36), ForeignKey("provider_endpoints.id", ondelete="SET NULL"), nullable=True
)
provider_api_key_id = Column(
String(36), ForeignKey("provider_api_keys.id", ondelete="SET NULL"), nullable=True
)
# Token统计
input_tokens = Column(Integer, default=0)
output_tokens = Column(Integer, default=0)
total_tokens = Column(Integer, default=0)
# 缓存相关 tokens (for Claude models)
cache_creation_input_tokens = Column(Integer, default=0)
cache_read_input_tokens = Column(Integer, default=0)
# 成本计算
input_cost_usd = Column(Float, default=0.0)
output_cost_usd = Column(Float, default=0.0)
cache_cost_usd = Column(Float, default=0.0) # 总缓存成本(兼容旧数据)
cache_creation_cost_usd = Column(Float, default=0.0) # 缓存创建成本
cache_read_cost_usd = Column(Float, default=0.0) # 缓存读取成本
request_cost_usd = Column(Float, default=0.0) # 按次计费成本
total_cost_usd = Column(Float, default=0.0)
# 真实成本计算(表面成本 x 倍率)
actual_input_cost_usd = Column(Float, default=0.0) # 真实输入成本
actual_output_cost_usd = Column(Float, default=0.0) # 真实输出成本
actual_cache_creation_cost_usd = Column(Float, default=0.0) # 真实缓存创建成本
actual_cache_read_cost_usd = Column(Float, default=0.0) # 真实缓存读取成本
actual_request_cost_usd = Column(Float, default=0.0) # 真实按次计费成本
actual_total_cost_usd = Column(Float, default=0.0) # 真实总成本
rate_multiplier = Column(Float, default=1.0) # 使用的倍率(来自 ProviderAPIKey
# 历史价格记录每1M tokens的美元价格记录请求时的实际价格
input_price_per_1m = Column(Float, nullable=True) # 输入单价
output_price_per_1m = Column(Float, nullable=True) # 输出单价
cache_creation_price_per_1m = Column(Float, nullable=True) # 缓存创建单价
cache_read_price_per_1m = Column(Float, nullable=True) # 缓存读取单价
price_per_request = Column(Float, nullable=True) # 按次计费单价(历史记录)
# 请求详情
request_type = Column(String(50)) # chat, completion, embedding等
api_format = Column(String(50), nullable=True) # API 格式: CLAUDE, OPENAI 等(用户请求格式)
endpoint_api_format = Column(String(50), nullable=True) # 端点原生 API 格式
has_format_conversion = Column(Boolean, nullable=True, default=False) # 是否发生了格式转换
is_stream = Column(Boolean, default=False) # 是否为流式请求
status_code = Column(Integer)
error_message = Column(Text, nullable=True)
error_category = Column(String(50), nullable=True, index=True)
response_time_ms = Column(Integer) # 总响应时间(毫秒)
first_byte_time_ms = Column(Integer, nullable=True) # 首字时间/TTFB毫秒
# 请求状态追踪
# pending: 请求开始处理中
# streaming: 流式响应进行中
# completed: 请求成功完成
# failed: 请求失败
# cancelled: 客户端主动断开连接
status = Column(String(20), default="completed", nullable=False, index=True)
# 结算状态(与 status 解耦)
# - pending: 等待结算(任务未完成 / 流式未结束)
# - settled: 已结算cost 已写入,可能 > 0 或 = 0
# - void: 作废(不收费,如任务未开始就取消)
billing_status = Column(String(20), default="settled", nullable=False, index=True)
finalized_at = Column(DateTime(timezone=True), nullable=True) # 结算完成时间(可选)
# 完整请求和响应记录
request_headers = Column(JSON, nullable=True) # 客户端请求头
request_body = Column(JSON, nullable=True) # 请求体7天内未压缩
provider_request_headers = Column(JSON, nullable=True) # 向提供商发送的请求头
response_headers = Column(JSON, nullable=True) # 提供商响应头
client_response_headers = Column(JSON, nullable=True) # 返回给客户端的响应头
response_body = Column(JSON, nullable=True) # 响应体7天内未压缩
# 压缩存储字段7天后自动压缩到这里
request_body_compressed = Column(LargeBinary, nullable=True) # gzip压缩的请求体
response_body_compressed = Column(LargeBinary, nullable=True) # gzip压缩的响应体
# 元数据
request_metadata = Column(JSON, nullable=True) # 存储额外信息
# 时间戳
created_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
nullable=False,
index=True,
)
# 关系
user = relationship("User", back_populates="usage_records")
api_key = relationship("ApiKey", back_populates="usage_records")
provider_obj = relationship("Provider") # 使用 provider_obj 避免与 provider 字段名冲突
provider_endpoint = relationship("ProviderEndpoint")
provider_api_key = relationship("ProviderAPIKey")
def get_request_body(self) -> Any:
"""获取请求体(自动解压)"""
if self.request_body is not None:
return self.request_body
if self.request_body_compressed is not None:
from src.utils.compression import decompress_json
return decompress_json(self.request_body_compressed)
return None
def get_response_body(self) -> Any:
"""获取响应体(自动解压)"""
if self.response_body is not None:
return self.response_body
if self.response_body_compressed is not None:
from src.utils.compression import decompress_json
return decompress_json(self.response_body_compressed)
return None
class RequestCandidate(Base):
"""请求候选记录 - 追踪所有候选(包括未使用的)"""
__tablename__ = "request_candidates"
# 主键
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# 关联字段
request_id = Column(String(100), nullable=False, index=True)
user_id = Column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=True)
api_key_id = Column(String(36), ForeignKey("api_keys.id", ondelete="CASCADE"), nullable=True)
# 候选信息
candidate_index = Column(Integer, nullable=False) # 候选序号从0开始
retry_index = Column(Integer, nullable=False, default=0) # 重试序号从0开始
provider_id = Column(String(36), ForeignKey("providers.id", ondelete="CASCADE"), nullable=True)
endpoint_id = Column(
String(36), ForeignKey("provider_endpoints.id", ondelete="CASCADE"), nullable=True
)
key_id = Column(
String(36), ForeignKey("provider_api_keys.id", ondelete="CASCADE"), nullable=True
)
# 状态信息
status = Column(
String(20), nullable=False
) # 'pending', 'streaming', 'success', 'failed', 'cancelled', 'skipped'
skip_reason = Column(Text, nullable=True) # 跳过/失败原因
is_cached = Column(Boolean, default=False) # 是否为缓存亲和性候选
# 执行结果信息(当 status = success/failed 时)
status_code = Column(Integer, nullable=True) # HTTP 状态码
error_type = Column(String(50), nullable=True) # 错误类型
error_message = Column(Text, nullable=True) # 错误消息
latency_ms = Column(Integer, nullable=True) # 延迟(毫秒)
concurrent_requests = Column(Integer, nullable=True) # 并发请求数
# 元数据
extra_data = Column(JSON, nullable=True)
required_capabilities = Column(JSON, nullable=True) # 请求实际需要的能力标签
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
started_at = Column(DateTime(timezone=True), nullable=True) # 开始执行时间
finished_at = Column(DateTime(timezone=True), nullable=True) # 完成时间
# 唯一约束和索引
__table_args__ = (
UniqueConstraint(
"request_id", "candidate_index", "retry_index", name="uq_request_candidate_with_retry"
),
Index("idx_request_candidates_request_id", "request_id"),
Index("idx_request_candidates_status", "status"),
Index("idx_request_candidates_provider_id", "provider_id"),
Index("idx_request_candidates_created_at", "created_at"),
)
# 关系
user = relationship("User")
api_key = relationship("ApiKey")
provider = relationship("Provider")
endpoint = relationship("ProviderEndpoint")
key = relationship("ProviderAPIKey")

522
src/models/user.py Normal file
View File

@@ -0,0 +1,522 @@
"""
用户相关数据库模型
包含: User, ApiKey, UserQuota, UserPreference, ManagementToken
"""
from __future__ import annotations
import hashlib
import secrets
import uuid
from datetime import datetime, timezone
import bcrypt
from sqlalchemy import (
JSON,
Boolean,
CheckConstraint,
Column,
DateTime,
Enum,
Float,
ForeignKey,
Index,
Integer,
String,
Text,
UniqueConstraint,
)
from sqlalchemy.orm import relationship
from src.config import config
from src.core.enums import AuthSource, UserRole
from ._base import Base
class User(Base):
"""用户模型"""
__tablename__ = "users"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
# OAuth 用户可能没有邮箱Postgres unique 允许多个 NULL
email = Column(String(255), unique=True, index=True, nullable=True)
# 注意:所有创建用户的入口必须显式写入 true/false禁止依赖默认值
email_verified = Column(Boolean, nullable=False)
username = Column(String(100), unique=True, index=True, nullable=False)
# OAuth 用户可能没有本地密码v1 仅做字段兼容)
password_hash = Column(String(255), nullable=True)
role = Column(
Enum(
UserRole,
name="userrole",
create_type=False,
values_callable=lambda x: [e.value for e in x],
),
default=UserRole.USER,
nullable=False,
)
auth_source = Column(
Enum(
AuthSource,
name="authsource",
create_type=False,
values_callable=lambda x: [e.value for e in x],
),
default=AuthSource.LOCAL,
nullable=False,
)
# LDAP 标识(仅 auth_source=ldap 时使用,用于在邮箱变更/用户名冲突时稳定关联本地账户)
ldap_dn = Column(String(512), nullable=True, index=True)
ldap_username = Column(String(255), nullable=True, index=True)
# 访问限制NULL 表示不限制,允许访问所有资源)
allowed_providers = Column(JSON, nullable=True) # 允许使用的提供商 ID 列表
allowed_api_formats = Column(JSON, nullable=True) # 允许使用的 API 格式列表
allowed_models = Column(JSON, nullable=True) # 允许使用的模型名称列表
# Key 能力配置
model_capability_settings = Column(JSON, nullable=True) # 用户针对特定模型的能力配置
# 示例: {"claude-sonnet-4-20250514": {"cache_1h": true}}
# 配额管理
quota_usd = Column(Float, nullable=True) # 美元配额(NULL 表示无限制)
used_usd = Column(Float, default=0.0) # 当前周期已使用美元
total_usd = Column(Float, default=0.0) # 累积消费总额
# 状态
is_active = Column(Boolean, default=True, nullable=False)
is_deleted = Column(Boolean, default=False, nullable=False)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
last_login_at = Column(DateTime(timezone=True), nullable=True)
# 关系 - CASCADE delete: 让数据库处理级联删除
api_keys = relationship("ApiKey", back_populates="user", cascade="all, delete-orphan")
management_tokens = relationship(
"ManagementToken", back_populates="user", cascade="all, delete-orphan"
)
preferences = relationship(
"UserPreference", back_populates="user", cascade="all, delete-orphan", passive_deletes=True
)
quotas = relationship(
"UserQuota", back_populates="user", cascade="all, delete-orphan", passive_deletes=True
)
announcement_reads = relationship(
"AnnouncementRead",
back_populates="user",
cascade="all, delete-orphan",
passive_deletes=True,
)
# 关系 - SET NULL: 保留历史记录,让数据库处理 SET NULL
usage_records = relationship("Usage", back_populates="user", passive_deletes=True)
authored_announcements = relationship(
"Announcement",
back_populates="author",
foreign_keys="Announcement.author_id",
passive_deletes=True,
)
audit_logs = relationship("AuditLog", back_populates="user", passive_deletes=True)
def set_password(self, password: str) -> None:
"""设置密码"""
self.password_hash = bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode(
"utf-8"
)
def verify_password(self, password: str) -> bool:
"""验证密码"""
if not self.password_hash:
return False
return bcrypt.checkpw(password.encode("utf-8"), self.password_hash.encode("utf-8"))
class ApiKey(Base):
"""API密钥模型"""
__tablename__ = "api_keys"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
user_id = Column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False)
key_hash = Column(String(64), unique=True, index=True, nullable=False) # API密钥的SHA256哈希
key_encrypted = Column(Text, nullable=True) # 加密后的完整密钥,用于查看
name = Column(String(100), nullable=True) # 密钥名称,便于用户管理
# 使用统计
total_requests = Column(Integer, default=0)
total_cost_usd = Column(Float, default=0.0)
# 余额管理(仅用于独立余额 Key
balance_used_usd = Column(Float, default=0.0) # 已使用余额USD用于统计
current_balance_usd = Column(Float, nullable=True) # 当前余额USDNULL 表示无限制
is_standalone = Column(
Boolean, default=False, nullable=False
) # 是否为独立余额 Key给非注册用户使用
# 访问限制NULL 表示不限制,允许访问所有资源)
allowed_providers = Column(JSON, nullable=True) # 允许使用的提供商 ID 列表
allowed_api_formats = Column(JSON, nullable=True) # 允许使用的 API 格式列表
allowed_models = Column(JSON, nullable=True) # 允许使用的模型名称列表
rate_limit = Column(Integer, default=None, nullable=True) # 每分钟请求限制None = 无限制
concurrent_limit = Column(Integer, default=5, nullable=True) # 并发请求限制
# Key 能力配置
force_capabilities = Column(JSON, nullable=True) # 强制开启的能力
# 示例: {"cache_1h": true} - 强制所有支持的模型都用 1h 缓存
# 状态
is_active = Column(Boolean, default=True, nullable=False)
is_locked = Column(Boolean, default=False, nullable=False) # 管理员锁定,用户无法使用/操作
last_used_at = Column(DateTime(timezone=True), nullable=True)
expires_at = Column(DateTime(timezone=True), nullable=True) # 过期时间
auto_delete_on_expiry = Column(Boolean, default=False, nullable=False) # 过期后是否自动删除
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
user = relationship("User", back_populates="api_keys")
usage_records = relationship("Usage", back_populates="api_key")
provider_mappings = relationship(
"ApiKeyProviderMapping", back_populates="api_key", cascade="all, delete-orphan"
)
@staticmethod
def generate_key() -> str:
"""生成API密钥使用加密安全的随机数生成器"""
import string
# 只使用字母和数字,避免特殊字符
alphabet = string.ascii_letters + string.digits
random_part = "".join(secrets.choice(alphabet) for _ in range(32))
return f"{config.api_key_prefix}-{random_part}"
@staticmethod
def hash_key(api_key: str) -> str:
"""对API密钥进行哈希"""
return hashlib.sha256(api_key.encode()).hexdigest()
def set_key(self, api_key: str) -> None:
"""
设置API密钥(用于测试和数据初始化)
Args:
api_key: 明文API密钥
注意: 此方法会设置 key_hash 和 key_encrypted
"""
from src.core.crypto import crypto_service
# 设置哈希(用于验证)
self.key_hash = self.hash_key(api_key)
# 设置加密的完整密钥(用于显示和管理)
self.key_encrypted = crypto_service.encrypt(api_key)
def verify_key(self, api_key: str) -> bool:
"""
验证API密钥是否匹配(用于测试)
Args:
api_key: 明文API密钥
Returns:
bool: 密钥是否匹配
"""
return self.key_hash == self.hash_key(api_key)
def get_display_key(self) -> str:
"""获取用于显示的脱敏密钥(前缀...后4位"""
from src.core.crypto import crypto_service
if self.key_encrypted:
try:
# 使用静默模式,避免在显示场景打印错误日志
full_key = crypto_service.decrypt(self.key_encrypted, silent=True)
# 格式sk-SpJ3y...sdf4
prefix = full_key[:10] if len(full_key) >= 10 else full_key[: len(full_key) // 2]
suffix = full_key[-4:] if len(full_key) >= 4 else ""
return f"{prefix}...{suffix}"
except Exception:
pass
# 降级:无法解密时返回占位符
return "sk-****"
class UserQuota(Base):
"""用户配额历史记录"""
__tablename__ = "user_quotas"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
user_id = Column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False)
# 配额类型
quota_type = Column(String(50), nullable=False) # monthly, daily, custom
# 配额值
quota_usd = Column(Float, nullable=False)
# 时间范围
period_start = Column(DateTime(timezone=True), nullable=False)
period_end = Column(DateTime(timezone=True), nullable=False)
# 使用情况
used_usd = Column(Float, default=0.0)
# 状态
is_active = Column(Boolean, default=True)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
user = relationship("User", back_populates="quotas")
class UserPreference(Base):
"""用户偏好设置表"""
__tablename__ = "user_preferences"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
user_id = Column(
String(36), ForeignKey("users.id", ondelete="CASCADE"), unique=True, nullable=False
)
# 个人信息
avatar_url = Column(String(500), nullable=True) # 头像URL
bio = Column(Text, nullable=True) # 个人简介
# 偏好设置
default_provider_id = Column(String(36), ForeignKey("providers.id"), nullable=True)
theme = Column(String(20), default="light") # light/dark/auto
language = Column(String(10), default="zh-CN")
timezone = Column(String(50), default="Asia/Shanghai")
# 通知设置
email_notifications = Column(Boolean, default=True)
usage_alerts = Column(Boolean, default=True)
announcement_notifications = Column(Boolean, default=True)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
user = relationship("User", back_populates="preferences")
default_provider = relationship("Provider")
class ManagementToken(Base):
"""Management Token 模型 - 用于程序化管理 API 调用"""
__tablename__ = "management_tokens"
# Token 格式常量
TOKEN_PREFIX = "ae_"
TOKEN_RANDOM_LENGTH = 40
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
user_id = Column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False)
# Token 信息
token_hash = Column(String(64), unique=True, index=True, nullable=False) # SHA256 哈希
token_prefix = Column(String(12), nullable=True) # Token 前缀用于显示(如 ae_xxxxxxxx
name = Column(String(100), nullable=False) # Token 名称
description = Column(Text, nullable=True) # 描述
# IP 白名单(可选)
allowed_ips = Column(JSON, nullable=True) # 允许的 IP 列表NULL = 不限制
# 格式: ["192.168.1.1", "10.0.0.0/24"]
# 有效期
expires_at = Column(DateTime(timezone=True), nullable=True) # NULL = 永不过期
# 使用统计
last_used_at = Column(DateTime(timezone=True), nullable=True)
last_used_ip = Column(String(45), nullable=True)
usage_count = Column(Integer, default=0) # 使用次数
# 状态
is_active = Column(Boolean, default=True, nullable=False)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
user = relationship("User", back_populates="management_tokens")
# 索引和约束
__table_args__ = (
Index("idx_management_tokens_user_id", "user_id"),
Index("idx_management_tokens_is_active", "is_active"),
UniqueConstraint("user_id", "name", name="uq_management_tokens_user_name"),
# IP 白名单必须为 NULL不限制或非空数组禁止空数组
# 注意JSON 类型的 NULL 可能被序列化为 JSON 'null',需要同时处理
CheckConstraint(
"allowed_ips IS NULL OR allowed_ips::text = 'null' OR json_array_length(allowed_ips) > 0",
name="check_allowed_ips_not_empty",
),
)
@staticmethod
def generate_token() -> str:
"""生成 Management Token使用加密安全的随机数"""
import string
alphabet = string.ascii_letters + string.digits
random_part = "".join(
secrets.choice(alphabet) for _ in range(ManagementToken.TOKEN_RANDOM_LENGTH)
)
return f"{ManagementToken.TOKEN_PREFIX}{random_part}"
@staticmethod
def hash_token(token: str) -> str:
"""对 Token 进行 SHA256 哈希
安全性说明(当前方案是安全的):
- Token 熵为 62^40约 2^238暴力破解在计算上不可行
- 结合速率限制(默认 30 次/分钟/IP在线攻击不可行
- 不需要盐值:盐值用于防止彩虹表攻击,但 Token 是高熵随机值,
不存在可预计算的"常见值",因此彩虹表攻击不适用
"""
return hashlib.sha256(token.encode()).hexdigest()
def set_token(self, token: str) -> None:
"""设置 Token只存储哈希和前缀用于显示"""
self.token_hash = self.hash_token(token)
# 存储前缀用于显示ae_ + 4 个字符,共 7 个字符)
self.token_prefix = token[:7] if len(token) > 7 else token
def get_display_token(self) -> str:
"""获取用于显示的脱敏 Token显示前缀 + 掩码)"""
if self.token_prefix:
return f"{self.token_prefix}...****"
return "ae_****"
def is_ip_allowed(self, client_ip: str) -> bool:
"""检查 IP 是否在白名单中
安全策略:
- None 或不设置表示不限制(允许所有 IP
- 非空列表表示只允许列表中的 IP
- 无效的白名单条目会被记录并跳过
- 无效的客户端 IP 直接拒绝
- 支持 IPv4 映射的 IPv6 地址规范化
"""
if self.allowed_ips is None:
return True # 未设置白名单,不限制
import ipaddress
from src.core.logger import logger
# 防御性检查:空列表应该在数据库层被拒绝,但这里再检查一次
if not self.allowed_ips:
logger.critical(f"Management Token {self.id} - allowed_ips 为空列表(违反数据库约束)")
return False # fail-safe
def normalize_ip(ip_str: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address | None:
"""规范化 IP 地址,将 IPv4 映射的 IPv6 转换为 IPv4"""
try:
ip = ipaddress.ip_address(ip_str)
if isinstance(ip, ipaddress.IPv6Address) and ip.ipv4_mapped:
return ip.ipv4_mapped
return ip
except ValueError:
return None
# 规范化客户端 IP
client = normalize_ip(client_ip)
if client is None:
logger.error(f"Management Token {self.id} - 拒绝无效的客户端 IP: {client_ip}")
return False
valid_entries = 0
for allowed in self.allowed_ips:
try:
if "/" in allowed:
# CIDR 格式
network = ipaddress.ip_network(allowed, strict=False)
valid_entries += 1
if client in network:
return True
else:
# 精确 IP
allowed_ip = normalize_ip(allowed)
if allowed_ip is None:
logger.error(f"Management Token {self.id} - 白名单包含无效条目: {allowed}")
continue
valid_entries += 1
if client == allowed_ip:
return True
except ValueError:
logger.error(f"Management Token {self.id} - 白名单包含无效条目: {allowed}")
continue
# 如果白名单全部无效,记录严重错误并拒绝
if valid_entries == 0:
logger.critical(f"Management Token {self.id} - 白名单全部无效,拒绝所有访问")
return False
@property
def is_expired(self) -> bool:
"""检查 Token 是否已过期(时区安全)"""
if not self.expires_at:
return False
expires = self.expires_at
if expires.tzinfo is None:
# 数据库中的时间应该有时区信息,如果没有则表示数据完整性问题
from src.core.logger import logger
logger.error(f"Management Token {self.id} expires_at 缺少时区信息(数据完整性问题)")
expires = expires.replace(tzinfo=timezone.utc)
return expires < datetime.now(timezone.utc)

View File

@@ -62,6 +62,7 @@ def update_user_agent_version(version: str) -> None:
# Backward compat: keep module-level constant in sync.
HTTP_USER_AGENT = f"antigravity/{_ua_version} {_PLATFORM_TAG}"
def parse_version_string(text: str) -> str | None:
"""从任意文本中提取 X.Y.Z 格式的版本号。"""
m = _VERSION_RE.search(text)
@@ -72,7 +73,9 @@ def parse_version_string(text: str) -> str | None:
URL_UNAVAILABLE_TTL_SECONDS = 300 # 5 分钟
# ============== Thinking Signature ==============
DUMMY_THOUGHT_SIGNATURE = "skip_thought_signature_validator"
# 统一从 core 层导入,避免多处定义
from src.core.api_format.conversion.constants import DUMMY_THOUGHT_SIGNATURE # noqa: E402
MIN_SIGNATURE_LENGTH = 50 # 与 Antigravity-Manager 对齐
# ============== Thinking Budget ==============

View File

@@ -3,22 +3,52 @@ from __future__ import annotations
import re
from collections.abc import Callable
from typing import Any
from uuid import uuid4
import httpx
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.error_utils import extract_error_message
from src.core.exceptions import (
ConcurrencyLimitError,
EmbeddedErrorException,
ProviderNotAvailableException,
ProxyNodeUnavailableError,
ThinkingSignatureException,
UpstreamClientException,
)
from src.core.logger import logger
from src.core.provider_types import ProviderType
from src.models.database import ApiKey
from src.models.database import (
ApiKey,
Provider,
ProviderAPIKey,
ProviderEndpoint,
RequestCandidate,
Usage,
User,
VideoTask,
)
from src.services.cache.aware_scheduler import (
CacheAwareScheduler,
get_cache_aware_scheduler,
)
from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import RetryPolicy, SkipPolicy
from src.services.candidate.recorder import CandidateRecorder
from src.services.candidate.resolver import CandidateResolver
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.orchestration.request_dispatcher import RequestDispatcher
from src.services.provider.format import normalize_endpoint_signature
from src.services.request.candidate import RequestCandidateService
from src.services.request.result import RequestMetadata
from src.services.system.config import SystemConfigService
from src.services.task.context import TaskMode
from src.services.task.exceptions import TaskNotFoundError
from src.services.task.protocol import AttemptKind, AttemptResult
from src.services.task.schema import ExecutionResult, TaskStatusResult
from src.services.usage.service import UsageService
_SENSITIVE_PATTERN = re.compile(
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
@@ -173,21 +203,9 @@ class TaskService:
- RequestDispatcher execution
- Error classification/rectify logic ported from the previous SYNC implementation
"""
from uuid import uuid4
from src.models.database import User
from src.services.cache.aware_scheduler import (
CacheAwareScheduler,
get_cache_aware_scheduler,
)
from src.services.candidate.resolver import CandidateResolver
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.orchestration.request_dispatcher import RequestDispatcher
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
from src.services.request.executor import RequestExecutor
from src.services.system.config import SystemConfigService
from src.services.usage.service import UsageService
if not request_id:
request_id = str(uuid4())
@@ -432,8 +450,6 @@ class TaskService:
if not error or not candidate:
return
from src.services.request.result import RequestMetadata
existing_metadata = getattr(error, "request_metadata", None)
if existing_metadata and getattr(existing_metadata, "api_format", None):
return
@@ -469,10 +485,6 @@ class TaskService:
last_error: Exception | None = None,
) -> None:
"""Raise a unified 'all candidates failed' exception."""
import httpx
from src.core.exceptions import ProviderNotAvailableException
logger.error(" [{}] 所有 {} 个组合均失败", request_id, max_attempts)
request_metadata = None
@@ -530,8 +542,6 @@ class TaskService:
extra_data: dict[str, Any],
) -> None:
"""Mark ThinkingSignatureException as failed for the candidate."""
from src.core.exceptions import ThinkingSignatureException
if not isinstance(error, ThinkingSignatureException):
return
@@ -560,7 +570,6 @@ class TaskService:
request_body_ref: dict[str, Any] | None,
) -> str:
"""Try to rectify thinking signature errors and request a retry."""
from src.core.exceptions import ThinkingSignatureException
from src.services.message.thinking_rectifier import ThinkingRectifier
if not isinstance(converted_error, ThinkingSignatureException):
@@ -697,17 +706,7 @@ class TaskService:
- "break": move to next candidate
- "raise": raise the underlying exception
"""
import httpx
from src.core.api_format.conversion.exceptions import FormatConversionError
from src.core.error_utils import extract_error_message
from src.core.exceptions import (
ConcurrencyLimitError,
EmbeddedErrorException,
ProxyNodeUnavailableError,
ThinkingSignatureException,
UpstreamClientException,
)
from src.services.proxy_node.resolver import resolve_effective_proxy, resolve_proxy_info
from src.services.request.executor import ExecutionError
@@ -980,20 +979,14 @@ class TaskService:
"""
from datetime import datetime, timezone
import httpx
from sqlalchemy import update
from src.models.database import RequestCandidate
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
from src.services.candidate.resolver import CandidateResolver
from src.services.candidate.submit import (
AllCandidatesFailedError,
SubmitOutcome,
UpstreamClientRequestError,
)
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.system.config import SystemConfigService
def _sanitize(message: str, max_length: int = 200) -> str:
if not message:
@@ -1131,8 +1124,6 @@ class TaskService:
continue
# 2. global switch (from database config)
from src.services.system.config import SystemConfigService
if not SystemConfigService.is_format_conversion_enabled(self.db):
skip_reason = "format_conversion_disabled"
candidate_info.update(
@@ -1398,8 +1389,6 @@ class TaskService:
- internal UUID (VideoTask.id)
- external operation id (VideoTask.short_id)
"""
from src.models.database import VideoTask
task = (
self.db.query(VideoTask)
.filter(VideoTask.id == task_id, VideoTask.user_id == user_id)
@@ -1498,9 +1487,7 @@ class TaskService:
)
from src.core.api_format.conversion.internal_video import VideoStatus
from src.core.crypto import crypto_service
from src.models.database import ProviderAPIKey, ProviderEndpoint
from src.services.provider.transport import build_provider_url
from src.services.usage.service import UsageService
try:
task = self._get_video_task_for_user(task_id, user_id=user_id)
@@ -1617,9 +1604,6 @@ class TaskService:
This keeps behavior compatible with the old Phase2 finalize logic.
"""
from src.models.database import ApiKey, Provider, User
from src.services.usage.service import UsageService
user_obj = self.db.query(User).filter(User.id == task.user_id).first()
api_key_obj = (
self.db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
@@ -1707,11 +1691,9 @@ class TaskService:
from datetime import datetime, timezone
from src.core.api_format.conversion.internal_video import VideoStatus
from src.models.database import Usage
from src.services.billing.dimension_collector_service import DimensionCollectorService
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
from src.services.billing.rule_service import BillingRuleService
from src.services.usage.service import UsageService
request_id = getattr(task, "request_id", None) or getattr(task, "id", None)
if not request_id:
@@ -1916,8 +1898,6 @@ class TaskService:
async def finalize(self, task_id: str) -> bool:
"""Finalize a task by internal id (best-effort)."""
from src.models.database import VideoTask
task = self.db.query(VideoTask).filter(VideoTask.id == task_id).first()
if not task:
return False

View File

@@ -0,0 +1,81 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
from sqlalchemy.orm import Session
from src.models.database import ApiKey, User
@dataclass
class UsageRecordParams:
"""用量记录参数数据类,用于在内部方法间传递数据"""
db: Session
user: User | None
api_key: ApiKey | None
provider: str
model: str
input_tokens: int
output_tokens: int
cache_creation_input_tokens: int
cache_read_input_tokens: int
request_type: str
api_format: str | None
endpoint_api_format: str | None # 端点原生 API 格式
has_format_conversion: bool # 是否发生了格式转换
is_stream: bool
response_time_ms: int | None
first_byte_time_ms: int | None
status_code: int
error_message: str | None
metadata: dict[str, Any] | None
request_headers: dict[str, Any] | None
request_body: Any | None
provider_request_headers: dict[str, Any] | None
response_headers: dict[str, Any] | None
client_response_headers: dict[str, Any] | None
response_body: Any | None
request_id: str
provider_id: str | None
provider_endpoint_id: str | None
provider_api_key_id: str | None
status: str
cache_ttl_minutes: int | None
use_tiered_pricing: bool
target_model: str | None
def __post_init__(self) -> None:
"""验证关键字段,确保数据完整性"""
# Token 数量不能为负数
if self.input_tokens < 0:
raise ValueError(f"input_tokens 不能为负数: {self.input_tokens}")
if self.output_tokens < 0:
raise ValueError(f"output_tokens 不能为负数: {self.output_tokens}")
if self.cache_creation_input_tokens < 0:
raise ValueError(
f"cache_creation_input_tokens 不能为负数: {self.cache_creation_input_tokens}"
)
if self.cache_read_input_tokens < 0:
raise ValueError(f"cache_read_input_tokens 不能为负数: {self.cache_read_input_tokens}")
# 响应时间不能为负数
if self.response_time_ms is not None and self.response_time_ms < 0:
raise ValueError(f"response_time_ms 不能为负数: {self.response_time_ms}")
if self.first_byte_time_ms is not None and self.first_byte_time_ms < 0:
raise ValueError(f"first_byte_time_ms 不能为负数: {self.first_byte_time_ms}")
# HTTP 状态码范围校验
if not (100 <= self.status_code <= 599):
raise ValueError(f"无效的 HTTP 状态码: {self.status_code}")
# 状态值校验
# - pending: 请求已创建,等待处理
# - streaming: 流式响应进行中
# - completed: 请求成功完成
# - failed: 请求失败(上游错误、超时等)
# - cancelled: 客户端主动断开连接
valid_statuses = {"pending", "streaming", "completed", "failed", "cancelled"}
if self.status not in valid_statuses:
raise ValueError(f"无效的状态值: {self.status},有效值: {valid_statuses}")

View File

@@ -0,0 +1,331 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.core.api_format.metadata import can_passthrough_endpoint
from src.core.api_format.signature import normalize_signature_key
from src.core.logger import logger
from src.models.database import RequestCandidate, Usage
class UsageActiveRequestsMixin:
"""活跃请求管理方法"""
@classmethod
def get_active_requests(
cls,
db: Session,
user_id: str | None = None,
limit: int = 50,
) -> list[Usage]:
"""
获取活跃的请求pending 或 streaming 状态)
Args:
db: 数据库会话
user_id: 用户ID可选用于过滤
limit: 最大返回数量
Returns:
活跃请求的 Usage 列表
"""
query = db.query(Usage).filter(Usage.status.in_(["pending", "streaming"]))
if user_id:
query = query.filter(Usage.user_id == user_id)
return query.order_by(Usage.created_at.desc()).limit(limit).all()
@classmethod
def cleanup_stale_pending_requests(
cls,
db: Session,
timeout_minutes: int = 10,
) -> int:
"""
清理超时的 pending/streaming 请求
将超过指定时间仍处于 pending 或 streaming 状态的请求标记为 failed。
这些请求可能是由于网络问题、服务重启或其他异常导致未能正常完成。
Args:
db: 数据库会话
timeout_minutes: 超时时间(分钟),默认 10 分钟
Returns:
清理的记录数
"""
cutoff_time = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
# 查找超时的请求
stale_requests = (
db.query(Usage)
.filter(
Usage.status.in_(["pending", "streaming"]),
Usage.created_at < cutoff_time,
)
.all()
)
count = 0
for usage in stale_requests:
old_status = usage.status
usage.status = "failed"
usage.error_message = f"请求超时: 状态 '{old_status}' 超过 {timeout_minutes} 分钟未完成"
usage.status_code = 504 # Gateway Timeout
count += 1
if count > 0:
db.commit()
logger.info(
f"清理超时请求: 将 {count} 条超过 {timeout_minutes} 分钟的 pending/streaming 请求标记为 failed"
)
return count
@classmethod
def get_stale_pending_count(
cls,
db: Session,
timeout_minutes: int = 10,
) -> int:
"""
获取超时的 pending/streaming 请求数量(用于监控)
Args:
db: 数据库会话
timeout_minutes: 超时时间(分钟)
Returns:
超时请求数量
"""
cutoff_time = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
return (
db.query(Usage)
.filter(
Usage.status.in_(["pending", "streaming"]),
Usage.created_at < cutoff_time,
)
.count()
)
@classmethod
def get_active_requests_status(
cls,
db: Session,
ids: list[str] | None = None,
user_id: str | None = None,
default_timeout_seconds: int = 300,
*,
include_admin_fields: bool = False,
) -> list[dict[str, Any]]:
"""
获取活跃请求状态(用于前端轮询),并自动清理超时的 pending/streaming 请求
与 get_active_requests 不同,此方法:
1. 返回轻量级的状态字典而非完整 Usage 对象
2. 自动检测并清理超时的 pending/streaming 请求
3. 支持按 ID 列表查询特定请求
Args:
db: 数据库会话
ids: 指定要查询的请求 ID 列表(可选)
user_id: 限制只查询该用户的请求(可选,用于普通用户接口)
default_timeout_seconds: 默认超时时间(秒),当端点未配置时使用
Returns:
请求状态列表
"""
now = datetime.now(timezone.utc)
# 构建基础查询
query = db.query(
Usage.id,
Usage.status,
Usage.input_tokens,
Usage.output_tokens,
Usage.cache_creation_input_tokens,
Usage.cache_read_input_tokens,
Usage.total_cost_usd,
Usage.actual_total_cost_usd,
Usage.rate_multiplier,
Usage.response_time_ms,
Usage.first_byte_time_ms, # 首字时间 (TTFB)
Usage.created_at,
Usage.provider_endpoint_id,
# API 格式 / 格式转换streaming 状态时已可确定)
Usage.api_format,
Usage.endpoint_api_format,
Usage.has_format_conversion,
# 模型映射streaming 时已可确定)
Usage.target_model,
)
# 管理员轮询:可附带 provider 与上游 key 名称(注意:不要在普通用户接口暴露上游 key 信息)
if include_admin_fields:
from src.models.database import ProviderAPIKey
query = query.add_columns(
Usage.provider_name,
ProviderAPIKey.name.label("api_key_name"),
).outerjoin(ProviderAPIKey, Usage.provider_api_key_id == ProviderAPIKey.id)
if ids:
query = query.filter(Usage.id.in_(ids))
if user_id:
query = query.filter(Usage.user_id == user_id)
else:
# 查询所有活跃请求
query = query.filter(Usage.status.in_(["pending", "streaming"]))
if user_id:
query = query.filter(Usage.user_id == user_id)
query = query.order_by(Usage.created_at.desc()).limit(50)
records = query.all()
# 检查超时的 pending/streaming 请求
# 收集可能超时的 usage_id 列表
timeout_candidates: list[str] = []
for r in records:
if r.status in ("pending", "streaming") and r.created_at:
# 使用全局配置的超时时间
timeout_seconds = default_timeout_seconds
# 处理时区:如果 created_at 没有时区信息,假定为 UTC
created_at = r.created_at
if created_at.tzinfo is None:
created_at = created_at.replace(tzinfo=timezone.utc)
elapsed = (now - created_at).total_seconds()
if elapsed > timeout_seconds:
# 需要获取 request_id 以便检查 RequestCandidate 表
# r.id 是 usage_id需要查询 request_id
timeout_candidates.append(r.id)
# 批量更新超时的请求(排除已有成功完成记录的请求)
timeout_ids = []
if timeout_candidates:
# 检查 RequestCandidate 表是否有成功完成的记录
# 如果流已经成功完成stream_completed: true不应该标记为超时
# 先获取这些 Usage 的 request_id
usage_request_ids = (
db.query(Usage.id, Usage.request_id).filter(Usage.id.in_(timeout_candidates)).all()
)
usage_id_to_request_id = {u.id: u.request_id for u in usage_request_ids}
request_id_to_usage_id = {u.request_id: u.id for u in usage_request_ids}
request_ids = list(request_id_to_usage_id.keys())
# 查询这些请求中已有成功完成记录的 request_id
# 包括两种情况:
# 1. status='success' 且 stream_completed=True正常完成
# 2. status='streaming' 且 status_code=200流传输中但 Provider 已返回 200可能是服务重启导致回调丢失
completed_usage_ids = set()
if request_ids:
from sqlalchemy import or_
candidates = (
db.query(
RequestCandidate.request_id,
RequestCandidate.status,
RequestCandidate.status_code,
RequestCandidate.extra_data,
)
.filter(
RequestCandidate.request_id.in_(request_ids),
or_(
RequestCandidate.status == "success",
# streaming 状态且 status_code=200说明 Provider 响应成功
# 但流传输可能因服务重启而中断
(RequestCandidate.status == "streaming")
& (RequestCandidate.status_code == 200),
),
)
.all()
)
for candidate in candidates:
extra_data = candidate.extra_data or {}
# 情况1status='success' 且 stream_completed=True
if candidate.status == "success" and extra_data.get("stream_completed", False):
usage_id = request_id_to_usage_id.get(candidate.request_id)
if usage_id:
completed_usage_ids.add(usage_id)
# 情况2status='streaming' 且 status_code=200
# 这表示 Provider 返回了 200但流传输可能因服务重启而未正常结束
# 此时应该恢复为 completed 而不是标记为 failed
elif candidate.status == "streaming" and candidate.status_code == 200:
usage_id = request_id_to_usage_id.get(candidate.request_id)
if usage_id:
completed_usage_ids.add(usage_id)
# 只对没有成功完成记录的请求标记超时
timeout_ids = [uid for uid in timeout_candidates if uid not in completed_usage_ids]
if timeout_ids:
db.query(Usage).filter(Usage.id.in_(timeout_ids)).update(
{"status": "failed", "error_message": "请求超时(服务器可能已重启)"},
synchronize_session=False,
)
db.commit()
# 对于已完成但状态未更新的请求,主动恢复状态为 completed
# 这处理了遥测回调丢失的情况(例如服务重启、后台任务未执行等)
if completed_usage_ids:
db.query(Usage).filter(Usage.id.in_(list(completed_usage_ids))).update(
{"status": "completed"},
synchronize_session=False,
)
db.commit()
logger.info(
f"[Usage] 恢复 {len(completed_usage_ids)} 个已完成请求的状态(遥测回调丢失)"
)
result: list[dict[str, Any]] = []
for r in records:
api_format = getattr(r, "api_format", None)
endpoint_api_format = getattr(r, "endpoint_api_format", None)
has_format_conversion = getattr(r, "has_format_conversion", None)
# 兼容历史数据:当 streaming 状态已拿到两个格式但 has_format_conversion 为空时,回填推断结果
if has_format_conversion is None and api_format and endpoint_api_format:
client_raw = str(api_format).strip()
endpoint_raw = str(endpoint_api_format).strip()
if ":" in client_raw and ":" in endpoint_raw:
client_fmt = normalize_signature_key(client_raw)
endpoint_fmt = normalize_signature_key(endpoint_raw)
has_format_conversion = not can_passthrough_endpoint(client_fmt, endpoint_fmt)
item: dict[str, Any] = {
"id": r.id,
"status": "failed" if r.id in timeout_ids else r.status,
"input_tokens": r.input_tokens,
"output_tokens": r.output_tokens,
"cache_creation_input_tokens": r.cache_creation_input_tokens,
"cache_read_input_tokens": r.cache_read_input_tokens,
"cost": float(r.total_cost_usd) if r.total_cost_usd else 0,
"actual_cost": (
float(r.actual_total_cost_usd) if r.actual_total_cost_usd is not None else None
),
"rate_multiplier": (
float(r.rate_multiplier) if r.rate_multiplier is not None else None
),
"response_time_ms": r.response_time_ms,
"first_byte_time_ms": r.first_byte_time_ms, # 首字时间 (TTFB)
}
if api_format:
item["api_format"] = api_format
if endpoint_api_format:
item["endpoint_api_format"] = endpoint_api_format
if has_format_conversion is not None:
item["has_format_conversion"] = bool(has_format_conversion)
# 模型映射streaming 时已可确定)
if r.target_model:
item["target_model"] = r.target_model
if include_admin_fields:
item["provider"] = r.provider_name
item["api_key_name"] = r.api_key_name
result.append(item)
return result

View File

@@ -0,0 +1,529 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.models.database import Usage, User
class UsageCacheAnalysisMixin:
"""缓存分析方法"""
@staticmethod
def analyze_cache_affinity_ttl(
db: Session,
user_id: str | None = None,
api_key_id: str | None = None,
hours: int = 168,
) -> dict[str, Any]:
"""
分析用户请求间隔分布,推荐合适的缓存亲和性 TTL
通过分析同一用户连续请求之间的时间间隔,判断用户的使用模式:
- 高频用户间隔短5 分钟 TTL 足够
- 中频用户15-30 分钟 TTL
- 低频用户(间隔长):需要 60 分钟 TTL
Args:
db: 数据库会话
user_id: 指定用户 ID可选为空则分析所有用户
api_key_id: 指定 API Key ID可选
hours: 分析最近多少小时的数据
Returns:
包含分析结果的字典
"""
from sqlalchemy import text
# 计算时间范围
start_date = datetime.now(timezone.utc) - timedelta(hours=hours)
# 构建 SQL 查询 - 使用窗口函数计算请求间隔
# 按 user_id 或 api_key_id 分组,计算同一组内连续请求的时间差
group_by_field = "api_key_id" if api_key_id else "user_id"
# 构建过滤条件
filter_clause = ""
if user_id or api_key_id:
filter_clause = f"AND {group_by_field} = :filter_id"
sql = text(f"""
WITH user_requests AS (
SELECT
{group_by_field} as group_id,
created_at,
LAG(created_at) OVER (
PARTITION BY {group_by_field}
ORDER BY created_at
) as prev_request_at
FROM usage
WHERE status = 'completed'
AND created_at > :start_date
AND {group_by_field} IS NOT NULL
{filter_clause}
),
intervals AS (
SELECT
group_id,
EXTRACT(EPOCH FROM (created_at - prev_request_at)) / 60.0 as interval_minutes
FROM user_requests
WHERE prev_request_at IS NOT NULL
),
user_stats AS (
SELECT
group_id,
COUNT(*) as request_count,
COUNT(*) FILTER (WHERE interval_minutes <= 5) as within_5min,
COUNT(*) FILTER (WHERE interval_minutes > 5 AND interval_minutes <= 15) as within_15min,
COUNT(*) FILTER (WHERE interval_minutes > 15 AND interval_minutes <= 30) as within_30min,
COUNT(*) FILTER (WHERE interval_minutes > 30 AND interval_minutes <= 60) as within_60min,
COUNT(*) FILTER (WHERE interval_minutes > 60) as over_60min,
PERCENTILE_CONT(0.5) WITHIN GROUP (ORDER BY interval_minutes) as median_interval,
PERCENTILE_CONT(0.75) WITHIN GROUP (ORDER BY interval_minutes) as p75_interval,
PERCENTILE_CONT(0.90) WITHIN GROUP (ORDER BY interval_minutes) as p90_interval,
AVG(interval_minutes) as avg_interval,
MIN(interval_minutes) as min_interval,
MAX(interval_minutes) as max_interval
FROM intervals
GROUP BY group_id
HAVING COUNT(*) >= 2
)
SELECT * FROM user_stats
ORDER BY request_count DESC
""")
params: dict[str, Any] = {
"start_date": start_date,
}
if user_id:
params["filter_id"] = user_id
elif api_key_id:
params["filter_id"] = api_key_id
result = db.execute(sql, params)
rows = result.fetchall()
# 收集所有 user_id 以便批量查询用户信息
group_ids = [row[0] for row in rows]
# 如果是按 user_id 分组,查询用户信息
user_info_map: dict[str, dict[str, str]] = {}
if group_by_field == "user_id" and group_ids:
users = db.query(User).filter(User.id.in_(group_ids)).all()
for user in users:
user_info_map[str(user.id)] = {
"username": str(user.username),
"email": str(user.email) if user.email else "",
}
# 处理结果
users_analysis = []
for row in rows:
# row 是一个 tuple按查询顺序访问
(
group_id,
request_count,
within_5min,
within_15min,
within_30min,
within_60min,
over_60min,
median_interval,
p75_interval,
p90_interval,
avg_interval,
min_interval,
max_interval,
) = row
# 计算推荐 TTL
recommended_ttl = UsageCacheAnalysisMixin._calculate_recommended_ttl(
p75_interval, p90_interval
)
# 获取用户信息
user_info = user_info_map.get(str(group_id), {})
# 计算各区间占比
total_intervals = request_count
users_analysis.append(
{
"group_id": group_id,
"username": user_info.get("username"),
"email": user_info.get("email"),
"request_count": request_count,
"interval_distribution": {
"within_5min": within_5min,
"within_15min": within_15min,
"within_30min": within_30min,
"within_60min": within_60min,
"over_60min": over_60min,
},
"interval_percentages": {
"within_5min": round(within_5min / total_intervals * 100, 1),
"within_15min": round(within_15min / total_intervals * 100, 1),
"within_30min": round(within_30min / total_intervals * 100, 1),
"within_60min": round(within_60min / total_intervals * 100, 1),
"over_60min": round(over_60min / total_intervals * 100, 1),
},
"percentiles": {
"p50": round(float(median_interval), 2) if median_interval else None,
"p75": round(float(p75_interval), 2) if p75_interval else None,
"p90": round(float(p90_interval), 2) if p90_interval else None,
},
"avg_interval_minutes": (
round(float(avg_interval), 2) if avg_interval else None
),
"min_interval_minutes": (
round(float(min_interval), 2) if min_interval else None
),
"max_interval_minutes": (
round(float(max_interval), 2) if max_interval else None
),
"recommended_ttl_minutes": recommended_ttl,
"recommendation_reason": UsageCacheAnalysisMixin._get_ttl_recommendation_reason(
recommended_ttl, p75_interval, p90_interval
),
}
)
# 汇总统计
ttl_distribution = {"5min": 0, "15min": 0, "30min": 0, "60min": 0}
for analysis in users_analysis:
ttl = analysis["recommended_ttl_minutes"]
if ttl <= 5:
ttl_distribution["5min"] += 1
elif ttl <= 15:
ttl_distribution["15min"] += 1
elif ttl <= 30:
ttl_distribution["30min"] += 1
else:
ttl_distribution["60min"] += 1
return {
"analysis_period_hours": hours,
"total_users_analyzed": len(users_analysis),
"ttl_distribution": ttl_distribution,
"users": users_analysis,
}
@staticmethod
def _calculate_recommended_ttl(
p75_interval: float | None,
p90_interval: float | None,
) -> int:
"""
根据请求间隔分布计算推荐的缓存 TTL
策略:
- 如果 90% 的请求间隔都在 5 分钟内 -> 5 分钟 TTL
- 如果 75% 的请求间隔在 15 分钟内 -> 15 分钟 TTL
- 如果 75% 的请求间隔在 30 分钟内 -> 30 分钟 TTL
- 否则 -> 60 分钟 TTL
"""
if p90_interval is None or p75_interval is None:
return 5 # 默认值
# 如果 90% 的间隔都在 5 分钟内
if p90_interval <= 5:
return 5
# 如果 75% 的间隔在 15 分钟内
if p75_interval <= 15:
return 15
# 如果 75% 的间隔在 30 分钟内
if p75_interval <= 30:
return 30
# 低频用户,需要更长的 TTL
return 60
@staticmethod
def _get_ttl_recommendation_reason(
ttl: int,
p75_interval: float | None,
p90_interval: float | None,
) -> str:
"""生成 TTL 推荐理由"""
if p75_interval is None or p90_interval is None:
return "数据不足,使用默认值"
if ttl == 5:
return f"高频用户90% 的请求间隔在 {p90_interval:.1f} 分钟内"
elif ttl == 15:
return f"中高频用户75% 的请求间隔在 {p75_interval:.1f} 分钟内"
elif ttl == 30:
return f"中频用户75% 的请求间隔在 {p75_interval:.1f} 分钟内"
else:
return f"低频用户75% 的请求间隔为 {p75_interval:.1f} 分钟,建议使用长 TTL"
@staticmethod
def get_cache_hit_analysis(
db: Session,
user_id: str | None = None,
api_key_id: str | None = None,
hours: int = 168,
) -> dict[str, Any]:
"""
分析缓存命中情况
Args:
db: 数据库会话
user_id: 指定用户 ID可选
api_key_id: 指定 API Key ID可选
hours: 分析最近多少小时的数据
Returns:
缓存命中分析结果
"""
start_date = datetime.now(timezone.utc) - timedelta(hours=hours)
# 基础查询
query = db.query(
func.count(Usage.id).label("total_requests"),
func.sum(Usage.input_tokens).label("total_input_tokens"),
func.sum(Usage.cache_read_input_tokens).label("total_cache_read_tokens"),
func.sum(Usage.cache_creation_input_tokens).label("total_cache_creation_tokens"),
func.sum(Usage.cache_read_cost_usd).label("total_cache_read_cost"),
func.sum(Usage.cache_creation_cost_usd).label("total_cache_creation_cost"),
).filter(
Usage.status == "completed",
Usage.created_at >= start_date,
)
if user_id:
query = query.filter(Usage.user_id == user_id)
if api_key_id:
query = query.filter(Usage.api_key_id == api_key_id)
result = query.first()
if result is None:
total_requests = 0
total_input_tokens = 0
total_cache_read_tokens = 0
total_cache_creation_tokens = 0
total_cache_read_cost = 0.0
total_cache_creation_cost = 0.0
else:
total_requests = result.total_requests or 0
total_input_tokens = result.total_input_tokens or 0
total_cache_read_tokens = result.total_cache_read_tokens or 0
total_cache_creation_tokens = result.total_cache_creation_tokens or 0
total_cache_read_cost = float(result.total_cache_read_cost or 0)
total_cache_creation_cost = float(result.total_cache_creation_cost or 0)
# 计算缓存命中率(按 token 数)
# 总输入上下文 = input_tokens + cache_read_tokens因为 input_tokens 不含 cache_read
# 或者如果 input_tokens 已经包含 cache_read则直接用 input_tokens
# 这里假设 cache_read_tokens 是额外的,命中率 = cache_read / (input + cache_read)
total_context_tokens = total_input_tokens + total_cache_read_tokens
cache_hit_rate = 0.0
if total_context_tokens > 0:
cache_hit_rate = total_cache_read_tokens / total_context_tokens * 100
# 计算节省的费用
# 缓存读取价格是正常输入价格的 10%,所以节省了 90%
# 节省 = cache_read_tokens * (正常价格 - 缓存价格) = cache_read_cost * 9
# 因为 cache_read_cost 是按 10% 价格算的,如果按 100% 算就是 10 倍
estimated_savings = total_cache_read_cost * 9 # 节省了 90%
# 统计有缓存命中的请求数
requests_with_cache_hit = db.query(func.count(Usage.id)).filter(
Usage.status == "completed",
Usage.created_at >= start_date,
Usage.cache_read_input_tokens > 0,
)
if user_id:
requests_with_cache_hit = requests_with_cache_hit.filter(Usage.user_id == user_id)
if api_key_id:
requests_with_cache_hit = requests_with_cache_hit.filter(Usage.api_key_id == api_key_id)
requests_with_cache_hit_count = int(requests_with_cache_hit.scalar() or 0)
return {
"analysis_period_hours": hours,
"total_requests": total_requests,
"requests_with_cache_hit": requests_with_cache_hit_count,
"request_cache_hit_rate": (
round(requests_with_cache_hit_count / total_requests * 100, 2)
if total_requests > 0
else 0
),
"total_input_tokens": total_input_tokens,
"total_cache_read_tokens": total_cache_read_tokens,
"total_cache_creation_tokens": total_cache_creation_tokens,
"token_cache_hit_rate": round(cache_hit_rate, 2),
"total_cache_read_cost_usd": round(total_cache_read_cost, 4),
"total_cache_creation_cost_usd": round(total_cache_creation_cost, 4),
"estimated_savings_usd": round(estimated_savings, 4),
}
@staticmethod
def get_interval_timeline(
db: Session,
hours: int = 24,
limit: int = 10000,
user_id: str | None = None,
include_user_info: bool = False,
) -> dict[str, Any]:
"""
获取请求间隔时间线数据,用于散点图展示
Args:
db: 数据库会话
hours: 分析最近多少小时的数据默认24小时
limit: 最大返回数据点数量默认10000
user_id: 指定用户 ID可选为空则返回所有用户
include_user_info: 是否包含用户信息(用于管理员多用户视图)
Returns:
包含时间线数据点的字典,每个数据点包含 model 字段用于按模型区分颜色
"""
from sqlalchemy import text
start_date = datetime.now(timezone.utc) - timedelta(hours=hours)
# 构建用户过滤条件
user_filter = "AND u.user_id = :user_id" if user_id else ""
# 根据是否需要用户信息选择不同的查询
if include_user_info and not user_id:
# 管理员视图:返回带用户信息的数据点
# 使用按比例采样,保持每个用户的数据量比例不变
sql = text(f"""
WITH request_intervals AS (
SELECT
u.created_at,
u.user_id,
u.model,
usr.username,
LAG(u.created_at) OVER (
PARTITION BY u.user_id
ORDER BY u.created_at
) as prev_request_at
FROM usage u
LEFT JOIN users usr ON u.user_id = usr.id
WHERE u.status = 'completed'
AND u.created_at > :start_date
AND u.user_id IS NOT NULL
{user_filter}
),
filtered_intervals AS (
SELECT
created_at,
user_id,
model,
username,
EXTRACT(EPOCH FROM (created_at - prev_request_at)) / 60.0 as interval_minutes,
ROW_NUMBER() OVER (PARTITION BY user_id ORDER BY created_at) as rn
FROM request_intervals
WHERE prev_request_at IS NOT NULL
AND EXTRACT(EPOCH FROM (created_at - prev_request_at)) / 60.0 <= 120
),
total_count AS (
SELECT COUNT(*) as cnt FROM filtered_intervals
),
user_totals AS (
SELECT user_id, COUNT(*) as user_cnt FROM filtered_intervals GROUP BY user_id
),
user_limits AS (
SELECT
ut.user_id,
CASE WHEN tc.cnt <= :limit THEN ut.user_cnt
ELSE GREATEST(CEIL(ut.user_cnt::float * :limit / tc.cnt), 1)::int
END as user_limit
FROM user_totals ut, total_count tc
)
SELECT
fi.created_at,
fi.user_id,
fi.model,
fi.username,
fi.interval_minutes
FROM filtered_intervals fi
JOIN user_limits ul ON fi.user_id = ul.user_id
WHERE fi.rn <= ul.user_limit
ORDER BY fi.created_at
""")
else:
# 普通视图:返回时间、间隔和模型信息
sql = text(f"""
WITH request_intervals AS (
SELECT
u.created_at,
u.user_id,
u.model,
LAG(u.created_at) OVER (
PARTITION BY u.user_id
ORDER BY u.created_at
) as prev_request_at
FROM usage u
WHERE u.status = 'completed'
AND u.created_at > :start_date
AND u.user_id IS NOT NULL
{user_filter}
)
SELECT
created_at,
model,
EXTRACT(EPOCH FROM (created_at - prev_request_at)) / 60.0 as interval_minutes
FROM request_intervals
WHERE prev_request_at IS NOT NULL
AND EXTRACT(EPOCH FROM (created_at - prev_request_at)) / 60.0 <= 120
ORDER BY created_at
LIMIT :limit
""")
params: dict[str, Any] = {"start_date": start_date, "limit": limit}
if user_id:
params["user_id"] = user_id
result = db.execute(sql, params)
rows = result.fetchall()
# 转换为时间线数据点
points = []
users_map: dict[str, str] = {} # user_id -> username
models_set: set = set() # 收集所有出现的模型
if include_user_info and not user_id:
for row in rows:
created_at, row_user_id, model, username, interval_minutes = row
point_data: dict[str, Any] = {
"x": created_at.isoformat(),
"y": round(float(interval_minutes), 2),
"user_id": str(row_user_id),
}
if model:
point_data["model"] = model
models_set.add(model)
points.append(point_data)
if row_user_id and username:
users_map[str(row_user_id)] = username
else:
for row in rows:
created_at, model, interval_minutes = row
point_data = {"x": created_at.isoformat(), "y": round(float(interval_minutes), 2)}
if model:
point_data["model"] = model
models_set.add(model)
points.append(point_data)
response: dict[str, Any] = {
"analysis_period_hours": hours,
"total_points": len(points),
"points": points,
}
if include_user_info and not user_id:
response["users"] = users_map
# 如果有模型信息,返回模型列表
if models_set:
response["models"] = sorted(models_set)
return response

View File

@@ -0,0 +1,524 @@
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, Usage, User
from src.services.system.config import SystemConfigService
class UsageLifecycleMixin:
"""使用记录生命周期管理方法"""
@classmethod
def begin_pending_usage(
cls,
db: Session,
request_id: str,
user: User | None,
api_key: ApiKey | None,
model: str,
*,
is_stream: bool = False,
request_type: str = "chat",
api_format: str | None = None,
request_headers: dict[str, Any] | None = None,
request_body: Any | None = None,
) -> Usage:
"""
创建或返回已有pending Usage 记录,但**不提交事务**。
适用场景:
- ApplicationService 在同一事务内创建 pending usage + task + candidates
- submit 幂等:重复调用同一 request_id 时返回已有记录
"""
existing = db.query(Usage).filter(Usage.request_id == request_id).first()
if existing:
return existing
# 根据配置决定是否记录请求详情
should_log_headers = SystemConfigService.should_log_headers(db)
should_log_body = SystemConfigService.should_log_body(db)
# 处理请求头
processed_request_headers = None
if should_log_headers and request_headers:
processed_request_headers = SystemConfigService.mask_sensitive_headers(
db, request_headers
)
# 处理请求体
processed_request_body = None
if should_log_body and request_body:
processed_request_body = SystemConfigService.truncate_body(
db, request_body, is_request=True
)
usage = Usage(
user_id=user.id if user else None,
api_key_id=api_key.id if api_key else None,
request_id=request_id,
provider_name="pending", # 尚未确定 provider
model=model,
input_tokens=0,
output_tokens=0,
total_tokens=0,
total_cost_usd=0.0,
request_type=request_type,
api_format=api_format,
is_stream=is_stream,
status="pending",
billing_status="pending",
request_headers=processed_request_headers,
request_body=processed_request_body,
)
db.add(usage)
db.flush()
return usage
@classmethod
def create_pending_usage(
cls,
db: Session,
request_id: str,
user: User | None,
api_key: ApiKey | None,
model: str,
is_stream: bool = False,
request_type: str = "chat",
api_format: str | None = None,
request_headers: dict[str, Any] | None = None,
request_body: Any | None = None,
) -> Usage:
"""
创建 pending 状态的使用记录(在请求开始时调用)
Args:
db: 数据库会话
request_id: 请求ID
user: 用户对象
api_key: API Key 对象
model: 模型名称
is_stream: 是否流式请求
api_format: API 格式
request_headers: 请求头
request_body: 请求体
Returns:
创建的 Usage 记录
"""
usage = cls.begin_pending_usage(
db,
request_id=request_id,
user=user,
api_key=api_key,
model=model,
is_stream=is_stream,
request_type=request_type,
api_format=api_format,
request_headers=request_headers,
request_body=request_body,
)
db.commit()
logger.debug("创建 pending 使用记录: request_id={}, model={}", request_id, model)
return usage
# ========== billing_status 并发幂等 finalize ==========
@classmethod
def finalize_settled(
cls,
db: Session,
request_id: str,
*,
total_cost_usd: float,
request_cost_usd: float | None = None,
status: str = "completed",
status_code: int = 200,
error_message: str | None = None,
response_time_ms: int | None = None,
billing_snapshot: dict[str, Any] | None = None,
extra_metadata: dict[str, Any] | None = None,
) -> bool:
"""
并发安全的幂等 finalizesettled
约定:
- 仅当 billing_status='pending' 时才会生效rowcount==1
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
cost = float(total_cost_usd)
request_cost = float(request_cost_usd) if request_cost_usd is not None else cost
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "pending",
)
.values(
billing_status="settled",
finalized_at=now,
total_cost_usd=cost,
request_cost_usd=request_cost,
status=status,
status_code=status_code,
error_message=error_message,
response_time_ms=response_time_ms,
)
)
if result.rowcount != 1:
return False
# 写入审计快照(只在本次 finalize 生效时执行)
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
if usage:
metadata = usage.request_metadata or {}
if billing_snapshot is not None:
metadata["billing_snapshot"] = billing_snapshot
if extra_metadata:
metadata.update(extra_metadata)
usage.request_metadata = cls._sanitize_request_metadata(metadata)
return True
@classmethod
def finalize_void(
cls,
db: Session,
request_id: str,
*,
reason: str | None = None,
status_code: int = 499,
) -> bool:
"""
并发安全的幂等 finalizevoid不收费
约定:
- 仅当 billing_status='pending' 时才会生效rowcount==1
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "pending",
)
.values(
billing_status="void",
finalized_at=now,
total_cost_usd=0.0,
request_cost_usd=0.0,
status="cancelled",
status_code=status_code,
error_message=reason,
response_time_ms=None,
)
)
return result.rowcount == 1
@classmethod
def finalize_submitted(
cls,
db: Session,
request_id: str,
*,
provider_name: str,
provider_id: str | None = None,
provider_endpoint_id: str | None = None,
provider_api_key_id: str | None = None,
response_time_ms: int | None = None,
status_code: int = 200,
endpoint_api_format: str | None = None,
provider_request_headers: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None,
response_body: Any | None = None,
) -> bool:
"""
异步任务提交成功时的幂等结算。
将 pending 使用记录标记为 settled费用暂时为 0。
后续轮询完成后通过 update_settled_billing 更新实际费用。
约定:
- 仅当 billing_status='pending' 时才会生效rowcount==1
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
# 处理响应头和响应体
should_log_headers = SystemConfigService.should_log_headers(db)
should_log_body = SystemConfigService.should_log_body(db)
processed_provider_headers = None
if should_log_headers and provider_request_headers:
processed_provider_headers = SystemConfigService.mask_sensitive_headers(
db, provider_request_headers
)
processed_response_headers = None
if should_log_headers and response_headers:
processed_response_headers = dict(response_headers)
processed_response_body = None
if should_log_body and response_body:
processed_response_body = SystemConfigService.truncate_body(
db, response_body, is_request=False
)
values: dict[str, Any] = {
"billing_status": "settled",
"finalized_at": now,
"total_cost_usd": 0.0,
"request_cost_usd": 0.0,
"status": "completed",
"status_code": status_code,
"response_time_ms": response_time_ms,
"provider_name": provider_name,
"provider_id": provider_id,
"provider_endpoint_id": provider_endpoint_id,
"provider_api_key_id": provider_api_key_id,
"endpoint_api_format": endpoint_api_format,
}
if processed_provider_headers is not None:
values["provider_request_headers"] = processed_provider_headers
if processed_response_headers is not None:
values["response_headers"] = processed_response_headers
if processed_response_body is not None:
values["response_body"] = processed_response_body
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "pending",
)
.values(**values)
)
return result.rowcount == 1
@classmethod
def update_settled_billing(
cls,
db: Session,
request_id: str,
*,
total_cost_usd: float,
request_cost_usd: float | None = None,
status: str = "completed",
status_code: int = 200,
error_message: str | None = None,
response_time_ms: int | None = None,
billing_snapshot: dict[str, Any] | None = None,
extra_metadata: dict[str, Any] | None = None,
) -> bool:
"""
更新已结算记录的计费信息(用于异步任务轮询完成后)。
与 finalize_settled 不同:
- finalize_settled: pending -> settled首次结算
- update_settled_billing: settled -> settled更新费用
约定:
- 仅当 billing_status='settled' 时才会生效
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
cost = float(total_cost_usd)
request_cost = float(request_cost_usd) if request_cost_usd is not None else cost
values: dict[str, Any] = {
"total_cost_usd": cost,
"request_cost_usd": request_cost,
"status": status,
"status_code": status_code,
}
if error_message is not None:
values["error_message"] = error_message
if response_time_ms is not None:
values["response_time_ms"] = response_time_ms
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "settled",
)
.values(**values)
)
if result.rowcount != 1:
return False
# 写入审计快照
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
if usage:
metadata = usage.request_metadata or {}
if billing_snapshot is not None:
metadata["billing_snapshot"] = billing_snapshot
if extra_metadata:
metadata.update(extra_metadata)
metadata["billing_updated_at"] = now.isoformat()
usage.request_metadata = cls._sanitize_request_metadata(metadata)
return True
@classmethod
def void_settled(
cls,
db: Session,
request_id: str,
*,
reason: str | None = None,
status_code: int = 499,
) -> bool:
"""
将已结算的记录作废(用于异步任务取消)。
与 finalize_void 不同:
- finalize_void: pending -> void未结算时作废
- void_settled: settled -> void已结算后取消费用归零
约定:
- 仅当 billing_status='settled' 时才会生效
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "settled",
)
.values(
billing_status="void",
finalized_at=now,
total_cost_usd=0.0,
request_cost_usd=0.0,
status="cancelled",
status_code=status_code,
error_message=reason,
)
)
return result.rowcount == 1
@classmethod
def update_usage_status(
cls,
db: Session,
request_id: str,
status: str,
error_message: str | None = None,
provider: str | None = None,
target_model: str | None = None,
first_byte_time_ms: int | None = None,
provider_id: str | None = None,
provider_endpoint_id: str | None = None,
provider_api_key_id: str | None = None,
api_format: str | None = None,
endpoint_api_format: str | None = None,
has_format_conversion: bool | None = None,
status_code: int | None = None,
) -> Usage | None:
"""
快速更新使用记录状态
Args:
db: 数据库会话
request_id: 请求ID
status: 新状态 (pending, streaming, completed, failed)
error_message: 错误消息(仅在 failed 状态时使用)
provider: 提供商名称可选streaming 状态时更新)
target_model: 映射后的目标模型名(可选)
first_byte_time_ms: 首字时间/TTFB可选streaming 状态时更新)
provider_id: Provider ID可选streaming 状态时更新)
provider_endpoint_id: Endpoint ID可选streaming 状态时更新)
provider_api_key_id: Provider API Key ID可选streaming 状态时更新)
api_format: API 格式(可选,用于获取按格式配置的倍率)
endpoint_api_format: 端点原生 API 格式(可选)
has_format_conversion: 是否发生了格式转换(可选)
status_code: HTTP 状态码(可选)
Returns:
更新后的 Usage 记录,如果未找到则返回 None
"""
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
if not usage:
logger.warning("未找到 request_id={} 的使用记录,无法更新状态", request_id)
return None
# 避免状态回退streaming 只能从 pending/streaming 进入
if status == "streaming" and usage.status not in ("pending", "streaming"):
logger.debug(
f"跳过 streaming 状态更新(避免回退): request_id={request_id}, "
f"{usage.status} -> {status}"
)
return usage
old_status = usage.status
usage.status = status
if error_message:
usage.error_message = error_message
if provider:
usage.provider_name = provider
elif status == "streaming" and usage.provider_name == "pending":
# 状态变为 streaming 但 provider_name 仍为 pending记录警告
logger.warning(
f"状态更新为 streaming 但 provider_name 为空: request_id={request_id}, "
f"当前 provider_name={usage.provider_name}"
)
if target_model:
usage.target_model = target_model
if first_byte_time_ms is not None:
usage.first_byte_time_ms = first_byte_time_ms
if provider_id is not None:
usage.provider_id = provider_id
if provider_endpoint_id is not None:
usage.provider_endpoint_id = provider_endpoint_id
if provider_api_key_id is not None:
usage.provider_api_key_id = provider_api_key_id
# 当设置 provider_api_key_id 时,同步获取并更新 rate_multiplier
# 这样前端在 streaming 状态就能显示倍率
rate_multiplier = cls._get_rate_multiplier_sync(
db, provider_api_key_id, api_format or usage.api_format
)
if rate_multiplier is not None:
usage.rate_multiplier = rate_multiplier
if endpoint_api_format is not None:
usage.endpoint_api_format = endpoint_api_format
if has_format_conversion is not None:
usage.has_format_conversion = has_format_conversion
if status_code is not None:
usage.status_code = status_code
# 结算状态:当请求进入终态时,将 billing_status 标记为 settled
# 注意:取消是否应 VOID/部分结算由更高层策略决定;这里默认终态均视为已结算。
if status in ("completed", "failed", "cancelled"):
if getattr(usage, "billing_status", None) == "pending":
usage.billing_status = "settled"
if getattr(usage, "finalized_at", None) is None:
usage.finalized_at = datetime.now(timezone.utc)
db.commit()
logger.debug("更新使用记录状态: request_id={}, {} -> {}", request_id, old_status, status)
return usage

View File

@@ -0,0 +1,176 @@
from __future__ import annotations
from sqlalchemy.orm import Session
from src.models.database import ProviderAPIKey
from src.services.model.cost import ModelCostService
class UsagePricingMixin:
"""定价相关方法"""
@classmethod
async def get_model_price_async(
cls, db: Session, provider: str, model: str
) -> tuple[float, float]:
"""异步获取模型价格输入价格输出价格每1M tokens
查找逻辑:
1. 直接通过 GlobalModel.name 匹配
2. 查找该 Provider 的 Model 实现并获取价格
3. 如果找不到则使用系统默认价格
"""
service = ModelCostService(db)
return await service.get_model_price_async(provider, model)
@classmethod
def get_model_price(cls, db: Session, provider: str, model: str) -> tuple[float, float]:
"""获取模型价格输入价格输出价格每1M tokens
查找逻辑:
1. 直接通过 GlobalModel.name 匹配
2. 查找该 Provider 的 Model 实现并获取价格
3. 如果找不到则使用系统默认价格
"""
service = ModelCostService(db)
return service.get_model_price(provider, model)
@classmethod
async def get_cache_prices_async(
cls, db: Session, provider: str, model: str, input_price: float
) -> tuple[float | None, float | None]:
"""异步获取模型缓存价格缓存创建价格缓存读取价格每1M tokens"""
service = ModelCostService(db)
return await service.get_cache_prices_async(provider, model, input_price)
@classmethod
def get_cache_prices(
cls, db: Session, provider: str, model: str, input_price: float
) -> tuple[float | None, float | None]:
"""获取模型缓存价格缓存创建价格缓存读取价格每1M tokens"""
service = ModelCostService(db)
return service.get_cache_prices(provider, model, input_price)
@classmethod
async def get_request_price_async(cls, db: Session, provider: str, model: str) -> float | None:
"""异步获取模型按次计费价格"""
service = ModelCostService(db)
return await service.get_request_price_async(provider, model)
@classmethod
def get_request_price(cls, db: Session, provider: str, model: str) -> float | None:
"""获取模型按次计费价格"""
service = ModelCostService(db)
return service.get_request_price(provider, model)
@staticmethod
def calculate_cost(
input_tokens: int,
output_tokens: int,
input_price_per_1m: float,
output_price_per_1m: float,
cache_creation_input_tokens: int = 0,
cache_read_input_tokens: int = 0,
cache_creation_price_per_1m: float | None = None,
cache_read_price_per_1m: float | None = None,
price_per_request: float | None = None,
) -> tuple[float, float, float, float, float, float, float]:
"""计算成本价格是每百万tokens- 固定价格模式
Returns:
Tuple of (input_cost, output_cost, cache_creation_cost,
cache_read_cost, cache_cost, request_cost, total_cost)
"""
return ModelCostService.compute_cost(
input_tokens=input_tokens,
output_tokens=output_tokens,
input_price_per_1m=input_price_per_1m,
output_price_per_1m=output_price_per_1m,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
cache_creation_price_per_1m=cache_creation_price_per_1m,
cache_read_price_per_1m=cache_read_price_per_1m,
price_per_request=price_per_request,
)
@classmethod
async def calculate_cost_with_strategy_async(
cls,
db: Session,
provider: str,
model: str,
input_tokens: int,
output_tokens: int,
cache_creation_input_tokens: int = 0,
cache_read_input_tokens: int = 0,
api_format: str | None = None,
cache_ttl_minutes: int | None = None,
) -> tuple[float, float, float, float, float, float, float, int | None]:
"""使用策略模式计算成本(支持阶梯计费)
根据 api_format 选择对应的计费策略,支持阶梯计费和 TTL 差异化。
Returns:
Tuple of (input_cost, output_cost, cache_creation_cost,
cache_read_cost, cache_cost, request_cost, total_cost, tier_index)
"""
service = ModelCostService(db)
return await service.compute_cost_with_strategy_async(
provider=provider,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
api_format=api_format,
cache_ttl_minutes=cache_ttl_minutes,
)
@classmethod
async def _get_rate_multiplier_and_free_tier(
cls,
db: Session,
provider_api_key_id: str | None,
provider_id: str | None,
api_format: str | None = None,
) -> tuple[float, bool]:
"""获取费率倍数和是否免费套餐(使用缓存)"""
from src.services.cache.provider_cache import ProviderCacheService
return await ProviderCacheService.get_rate_multiplier_and_free_tier(
db, provider_api_key_id, provider_id, api_format
)
@staticmethod
def _get_rate_multiplier_sync(
db: Session,
provider_api_key_id: str,
api_format: str | None = None,
) -> float | None:
"""
同步获取 ProviderAPIKey 的 rate_multiplier
Args:
db: 数据库会话
provider_api_key_id: ProviderAPIKey ID
api_format: API 格式(可选),如 "CLAUDE""OPENAI"
Returns:
rate_multiplier 或 None
"""
from src.services.cache.provider_cache import ProviderCacheService
provider_key = (
db.query(ProviderAPIKey.rate_multipliers)
.filter(ProviderAPIKey.id == provider_api_key_id)
.first()
)
if not provider_key:
return None
return ProviderCacheService.compute_rate_multiplier(
provider_key.rate_multipliers, api_format
)

531
src/services/usage/query.py Normal file
View File

@@ -0,0 +1,531 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, Usage, User, UserRole
class UsageQueryMixin:
"""查询/统计相关方法"""
# 热力图缓存键前缀(依赖 TTL 自动过期,用户角色变更时主动清除)
HEATMAP_CACHE_KEY_PREFIX = "activity_heatmap"
@classmethod
def _get_heatmap_cache_key(cls, user_id: str | None, include_actual_cost: bool) -> str:
"""生成热力图缓存键"""
cost_suffix = "with_cost" if include_actual_cost else "no_cost"
if user_id:
return f"{cls.HEATMAP_CACHE_KEY_PREFIX}:user:{user_id}:{cost_suffix}"
else:
return f"{cls.HEATMAP_CACHE_KEY_PREFIX}:admin:all:{cost_suffix}"
@classmethod
async def clear_user_heatmap_cache(cls, user_id: str) -> None:
"""
清除用户的热力图缓存(用户角色变更时调用)
Args:
user_id: 用户ID
"""
from src.clients.redis_client import get_redis_client
redis_client = await get_redis_client(require_redis=False)
if not redis_client:
return
# 清除该用户的所有热力图缓存with_cost 和 no_cost
keys_to_delete = [
cls._get_heatmap_cache_key(user_id, include_actual_cost=True),
cls._get_heatmap_cache_key(user_id, include_actual_cost=False),
]
for key in keys_to_delete:
try:
await redis_client.delete(key)
logger.debug("已清除热力图缓存: {}", key)
except Exception as e:
logger.warning("清除热力图缓存失败: {}, error={}", key, e)
@classmethod
async def get_cached_heatmap(
cls,
db: Session,
user_id: str | None = None,
include_actual_cost: bool = False,
) -> dict[str, Any]:
"""
获取带缓存的热力图数据
缓存策略:
- TTL: 10分钟CacheTTL.ACTIVITY_HEATMAP = 600
- 仅依赖 TTL 自动过期,新使用记录最多延迟 10 分钟出现
- 用户角色变更时通过 clear_user_heatmap_cache() 主动清除
Args:
db: 数据库会话
user_id: 用户IDNone 表示获取全局热力图(管理员)
include_actual_cost: 是否包含实际成本
Returns:
热力图数据字典
"""
import json
from src.clients.redis_client import get_redis_client
from src.config.constants import CacheTTL
cache_key = cls._get_heatmap_cache_key(user_id, include_actual_cost)
cache_ttl = CacheTTL.ACTIVITY_HEATMAP
redis_client = await get_redis_client(require_redis=False)
# 尝试从缓存获取
if redis_client:
try:
cached = await redis_client.get(cache_key)
if cached:
try:
return json.loads(cached) # type: ignore[no-any-return]
except json.JSONDecodeError as e:
logger.warning(
"热力图缓存解析失败,删除损坏缓存: {}, error={}", cache_key, e
)
try:
await redis_client.delete(cache_key)
except Exception:
pass
except Exception as e:
logger.error("读取热力图缓存出错: {}, error={}", cache_key, e)
# 从数据库查询
result = cls.get_daily_activity(
db=db,
user_id=user_id,
window_days=365,
include_actual_cost=include_actual_cost,
)
# 保存到缓存(失败不影响返回结果)
if redis_client:
try:
await redis_client.setex(
cache_key,
cache_ttl,
json.dumps(result, ensure_ascii=False, default=str),
)
except Exception as e:
logger.warning("保存热力图缓存失败: {}, error={}", cache_key, e)
return result
@staticmethod
def check_user_quota(
db: Session,
user: User,
estimated_tokens: int = 0,
estimated_cost: float = 0,
api_key: ApiKey | None = None,
) -> tuple[bool, str]:
"""检查用户配额或独立Key余额
Args:
db: 数据库会话
user: 用户对象
estimated_tokens: 预估token数
estimated_cost: 预估费用
api_key: API Key对象用于检查独立余额Key
Returns:
(是否通过, 消息)
"""
# 如果是独立余额Key检查Key的余额而不是用户配额
if api_key and api_key.is_standalone:
# 导入 ApiKeyService 以使用统一的余额计算方法
from src.services.user.apikey import ApiKeyService
# NULL 表示无限制
if api_key.current_balance_usd is None:
return True, "OK"
# 使用统一的余额计算方法
remaining_balance = ApiKeyService.get_remaining_balance(api_key)
if remaining_balance is None:
return True, "OK"
# 检查余额是否充足
if remaining_balance < estimated_cost:
return (
False,
f"Key余额不足剩余: ${remaining_balance:.2f},需要: ${estimated_cost:.2f}",
)
return True, "OK"
# 普通Key检查用户配额
# 管理员无限制
if user.role == UserRole.ADMIN:
return True, "OK"
# NULL 表示无限制
if user.quota_usd is None:
return True, "OK"
# 有配额限制,检查是否超额
used_usd = float(user.used_usd or 0)
quota_usd = float(user.quota_usd)
if used_usd + estimated_cost > quota_usd:
remaining = quota_usd - used_usd
return False, f"配额不足(剩余: ${remaining:.2f}"
return True, "OK"
@staticmethod
def get_usage_summary(
db: Session,
user_id: str | None = None,
api_key_id: str | None = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
group_by: str = "day", # day, week, month
) -> list[dict[str, Any]]:
"""获取使用汇总"""
query = db.query(Usage)
# 过滤掉 pending/streaming 状态的请求(尚未完成的请求不应计入统计)
query = query.filter(Usage.status.notin_(["pending", "streaming"]))
if user_id:
query = query.filter(Usage.user_id == user_id)
if api_key_id:
query = query.filter(Usage.api_key_id == api_key_id)
if start_date:
query = query.filter(Usage.created_at >= start_date)
if end_date:
query = query.filter(Usage.created_at < end_date)
# 使用跨数据库兼容的日期函数
from src.utils.database_helpers import date_trunc_portable
# 检测数据库方言
bind = db.bind
dialect = bind.dialect.name if bind is not None else "sqlite"
# 根据分组类型选择日期函数(兼容多种数据库)
if group_by == "day":
date_func = date_trunc_portable(dialect, "day", Usage.created_at)
elif group_by == "week":
date_func = date_trunc_portable(dialect, "week", Usage.created_at)
elif group_by == "month":
date_func = date_trunc_portable(dialect, "month", Usage.created_at)
else:
# 默认按天分组
date_func = date_trunc_portable(dialect, "day", Usage.created_at)
# 汇总查询
summary = db.query(
date_func.label("period"),
Usage.provider_name,
Usage.model,
func.count(Usage.id).label("requests"),
func.sum(Usage.input_tokens).label("input_tokens"),
func.sum(Usage.output_tokens).label("output_tokens"),
func.sum(Usage.total_tokens).label("total_tokens"),
func.sum(Usage.total_cost_usd).label("total_cost_usd"),
func.avg(Usage.response_time_ms).label("avg_response_time"),
)
# 过滤掉 pending/streaming 状态的请求(与上方明细查询一致)
summary = summary.filter(Usage.status.notin_(["pending", "streaming"]))
if user_id:
summary = summary.filter(Usage.user_id == user_id)
if api_key_id:
summary = summary.filter(Usage.api_key_id == api_key_id)
if start_date:
summary = summary.filter(Usage.created_at >= start_date)
if end_date:
summary = summary.filter(Usage.created_at < end_date)
summary = summary.group_by(date_func, Usage.provider_name, Usage.model).all()
return [
{
"period": row.period,
"provider": row.provider_name,
"model": row.model,
"requests": row.requests,
"input_tokens": row.input_tokens,
"output_tokens": row.output_tokens,
"total_tokens": row.total_tokens,
"total_cost_usd": float(row.total_cost_usd),
"avg_response_time_ms": (
float(row.avg_response_time) if row.avg_response_time else 0
),
}
for row in summary
]
@staticmethod
def get_daily_activity(
db: Session,
user_id: str | None = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
window_days: int = 365,
include_actual_cost: bool = False,
) -> dict[str, Any]:
"""按天统计请求活跃度,用于渲染热力图。
优化策略:
- 历史数据从预计算的 StatsDaily/StatsUserDaily 表读取
- 只有"今天"的数据才实时查询 Usage 表
"""
def ensure_timezone(value: datetime) -> datetime:
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
# 如果调用方未指定时间范围,则默认统计最近 window_days 天
now = datetime.now(timezone.utc)
end_dt = ensure_timezone(end_date) if end_date else now
start_dt = (
ensure_timezone(start_date) if start_date else end_dt - timedelta(days=window_days - 1)
)
# 对齐到自然日的开始/结束
start_dt = datetime.combine(start_dt.date(), datetime.min.time(), tzinfo=timezone.utc)
end_dt = datetime.combine(end_dt.date(), datetime.max.time(), tzinfo=timezone.utc)
today = now.date()
today_start_dt = datetime.combine(today, datetime.min.time(), tzinfo=timezone.utc)
aggregated: dict[str, dict[str, Any]] = {}
# 1. 从预计算表读取历史数据(不包括今天)
if user_id:
from src.models.database import StatsUserDaily
hist_query = db.query(StatsUserDaily).filter(
StatsUserDaily.user_id == user_id,
StatsUserDaily.date >= start_dt,
StatsUserDaily.date < today_start_dt,
)
for row in hist_query.all():
key = (
row.date.date().isoformat()
if isinstance(row.date, datetime)
else str(row.date)[:10]
)
aggregated[key] = {
"requests": row.total_requests or 0,
"total_tokens": (
(row.input_tokens or 0)
+ (row.output_tokens or 0)
+ (row.cache_creation_tokens or 0)
+ (row.cache_read_tokens or 0)
),
"total_cost_usd": float(row.total_cost or 0.0),
}
# StatsUserDaily 没有 actual_total_cost 字段,用户视图不需要倍率成本
else:
from src.models.database import StatsDaily
hist_query = db.query(StatsDaily).filter(
StatsDaily.date >= start_dt,
StatsDaily.date < today_start_dt,
)
for row in hist_query.all():
key = (
row.date.date().isoformat()
if isinstance(row.date, datetime)
else str(row.date)[:10]
)
aggregated[key] = {
"requests": row.total_requests or 0,
"total_tokens": (
(row.input_tokens or 0)
+ (row.output_tokens or 0)
+ (row.cache_creation_tokens or 0)
+ (row.cache_read_tokens or 0)
),
"total_cost_usd": float(row.total_cost or 0.0),
}
if include_actual_cost:
aggregated[key]["actual_total_cost_usd"] = float(
row.actual_total_cost or 0.0 # type: ignore[attr-defined]
)
# 2. 实时查询今天的数据(如果在查询范围内)
if today >= start_dt.date() and today <= end_dt.date():
today_start = datetime.combine(today, datetime.min.time(), tzinfo=timezone.utc)
today_end = datetime.combine(today, datetime.max.time(), tzinfo=timezone.utc)
if include_actual_cost:
today_query = db.query(
func.count(Usage.id).label("requests"),
func.sum(Usage.total_tokens).label("total_tokens"),
func.sum(Usage.total_cost_usd).label("total_cost_usd"),
func.sum(Usage.actual_total_cost_usd).label("actual_total_cost_usd"),
).filter(
Usage.created_at >= today_start,
Usage.created_at <= today_end,
)
else:
today_query = db.query(
func.count(Usage.id).label("requests"),
func.sum(Usage.total_tokens).label("total_tokens"),
func.sum(Usage.total_cost_usd).label("total_cost_usd"),
).filter(
Usage.created_at >= today_start,
Usage.created_at <= today_end,
)
if user_id:
today_query = today_query.filter(Usage.user_id == user_id)
today_row = today_query.first()
if today_row and today_row.requests:
aggregated[today.isoformat()] = {
"requests": int(today_row.requests or 0),
"total_tokens": int(today_row.total_tokens or 0),
"total_cost_usd": float(today_row.total_cost_usd or 0.0),
}
if include_actual_cost:
aggregated[today.isoformat()]["actual_total_cost_usd"] = float(
today_row.actual_total_cost_usd or 0.0
)
# 3. 构建返回结果
days: list[dict[str, Any]] = []
cursor = start_dt.date()
end_date_only = end_dt.date()
max_requests = 0
while cursor <= end_date_only:
iso_date = cursor.isoformat()
stats = aggregated.get(iso_date, {})
requests = stats.get("requests", 0)
total_tokens = stats.get("total_tokens", 0)
total_cost = stats.get("total_cost_usd", 0.0)
entry: dict[str, Any] = {
"date": iso_date,
"requests": requests,
"total_tokens": total_tokens,
"total_cost": total_cost,
}
if include_actual_cost:
entry["actual_total_cost"] = stats.get("actual_total_cost_usd", 0.0)
days.append(entry)
max_requests = max(max_requests, requests)
cursor += timedelta(days=1)
return {
"start_date": start_dt.date().isoformat(),
"end_date": end_dt.date().isoformat(),
"total_days": len(days),
"max_requests": max_requests,
"days": days,
}
@staticmethod
def get_top_users(
db: Session,
limit: int = 10,
start_date: datetime | None = None,
end_date: datetime | None = None,
order_by: str = "cost", # cost, tokens, requests
) -> list[dict[str, Any]]:
"""获取使用量最高的用户"""
query = (
db.query(
User.id,
User.email,
User.username,
func.count(Usage.id).label("requests"),
func.sum(Usage.total_tokens).label("tokens"),
func.sum(Usage.total_cost_usd).label("cost_usd"),
)
.join(Usage, User.id == Usage.user_id)
.filter(Usage.user_id.isnot(None))
)
if start_date:
query = query.filter(Usage.created_at >= start_date)
if end_date:
query = query.filter(Usage.created_at <= end_date)
query = query.group_by(User.id, User.email, User.username)
# 排序
if order_by == "cost":
query = query.order_by(func.sum(Usage.total_cost_usd).desc())
elif order_by == "tokens":
query = query.order_by(func.sum(Usage.total_tokens).desc())
else:
query = query.order_by(func.count(Usage.id).desc())
results = query.limit(limit).all()
return [
{
"user_id": row.id,
"email": row.email,
"username": row.username,
"requests": row.requests,
"tokens": row.tokens,
"cost_usd": float(row.cost_usd),
}
for row in results
]
@staticmethod
def cleanup_old_usage_records(
db: Session, days_to_keep: int = 90, batch_size: int = 1000
) -> int:
"""清理旧的使用记录(分批删除避免长事务锁定)
Args:
db: 数据库会话
days_to_keep: 保留天数,默认 90 天
batch_size: 每批删除数量,默认 1000 条
Returns:
删除的总记录数
"""
cutoff_date = datetime.now(timezone.utc) - timedelta(days=days_to_keep)
total_deleted = 0
while True:
# 查询待删除的 ID使用新索引 idx_usage_user_created
batch_ids = (
db.query(Usage.id).filter(Usage.created_at < cutoff_date).limit(batch_size).all()
)
if not batch_ids:
break
# 批量删除
deleted_count = (
db.query(Usage)
.filter(Usage.id.in_([row.id for row in batch_ids]))
.delete(synchronize_session=False)
)
db.commit()
total_deleted += deleted_count
logger.debug("清理使用记录: 本批删除 {}", deleted_count)
logger.info("清理使用记录: 共删除 {} 条超过 {} 天的记录", total_deleted, days_to_keep)
return total_deleted

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff