mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
Merge remote-tracking branch 'origin/master' into dev
# Conflicts: # src/api/handlers/base/request_builder.py # src/api/handlers/openai_cli/adapter.py # src/services/system/maintenance_scheduler.py
This commit is contained in:
@@ -126,16 +126,7 @@ class MessageTelemetry:
|
||||
# Provider 响应元数据(如 Gemini 的 modelVersion)
|
||||
response_metadata: dict[str, Any] | None = None,
|
||||
) -> float:
|
||||
total_cost = await self.calculate_cost(
|
||||
provider,
|
||||
model,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_creation_tokens=cache_creation_tokens,
|
||||
cache_read_tokens=cache_read_tokens,
|
||||
)
|
||||
|
||||
await UsageService.record_usage(
|
||||
usage = await UsageService.record_usage(
|
||||
db=self.db,
|
||||
user=self.user,
|
||||
api_key=self.api_key,
|
||||
@@ -170,6 +161,8 @@ class MessageTelemetry:
|
||||
metadata=response_metadata,
|
||||
)
|
||||
|
||||
total_cost = float(getattr(usage, "total_cost_usd", 0.0) or 0.0)
|
||||
|
||||
if self.user and self.api_key:
|
||||
audit_service.log_api_request(
|
||||
db=self.db,
|
||||
@@ -181,8 +174,8 @@ class MessageTelemetry:
|
||||
success=True,
|
||||
ip_address=self.client_ip,
|
||||
status_code=status_code,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
input_tokens=getattr(usage, "input_tokens", input_tokens),
|
||||
output_tokens=getattr(usage, "output_tokens", output_tokens),
|
||||
cost_usd=total_cost,
|
||||
)
|
||||
|
||||
|
||||
@@ -630,6 +630,9 @@ class ChatAdapterBase(ApiAdapter):
|
||||
api_key: str,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 端点规则参数
|
||||
body_rules: list[dict[str, Any]] | None = None,
|
||||
header_rules: list[dict[str, Any]] | None = None,
|
||||
# 用量计算参数(现在强制记录)
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
@@ -647,6 +650,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
api_key: API 密钥(已解密)
|
||||
request_data: 请求数据
|
||||
extra_headers: 端点配置的额外请求头
|
||||
body_rules: 请求体规则(在格式转换后应用)
|
||||
header_rules: 请求头规则(在请求头构建后应用)
|
||||
db: 数据库会话
|
||||
user: 用户对象
|
||||
provider_name: 提供商名称
|
||||
@@ -658,12 +663,30 @@ class ChatAdapterBase(ApiAdapter):
|
||||
测试响应数据
|
||||
"""
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
from src.api.handlers.base.request_builder import apply_body_rules
|
||||
from src.core.api_format.headers import HeaderBuilder
|
||||
|
||||
# 使用子类配置方法构建请求组件
|
||||
url = cls.build_endpoint_url(base_url)
|
||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||
body = cls.build_request_body(request_data)
|
||||
|
||||
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
|
||||
if body_rules:
|
||||
body = apply_body_rules(body, body_rules)
|
||||
|
||||
# 应用请求头规则(在请求头构建后应用)
|
||||
if header_rules:
|
||||
# 获取认证头名称,防止被规则覆盖
|
||||
from src.core.api_format import get_auth_config_for_endpoint
|
||||
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
|
||||
header_builder = HeaderBuilder()
|
||||
header_builder.add_many(headers)
|
||||
header_builder.apply_rules(header_rules, protected_keys)
|
||||
headers = header_builder.build()
|
||||
|
||||
# 使用通用的endpoint checker执行请求
|
||||
return await run_endpoint_check(
|
||||
client=client,
|
||||
|
||||
@@ -598,6 +598,9 @@ class CliAdapterBase(ApiAdapter):
|
||||
api_key: str,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 端点规则参数
|
||||
body_rules: list[dict[str, Any]] | None = None,
|
||||
header_rules: list[dict[str, Any]] | None = None,
|
||||
# 用量计算参数
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
@@ -622,6 +625,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
api_key: API 密钥(已解密)
|
||||
request_data: 请求数据
|
||||
extra_headers: 端点配置的额外请求头
|
||||
body_rules: 请求体规则(在格式转换后应用)
|
||||
header_rules: 请求头规则(在请求头构建后应用)
|
||||
db: 数据库会话
|
||||
user: 用户对象
|
||||
provider_name: 提供商名称
|
||||
@@ -633,6 +638,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
测试响应数据
|
||||
"""
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
from src.api.handlers.base.request_builder import apply_body_rules
|
||||
from src.core.api_format.headers import HeaderBuilder
|
||||
|
||||
# 构建请求组件
|
||||
url = cls.build_endpoint_url(base_url, request_data, model_name)
|
||||
@@ -646,6 +653,22 @@ class CliAdapterBase(ApiAdapter):
|
||||
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
|
||||
body = cls.build_request_body(request_data, base_url=base_url)
|
||||
|
||||
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
|
||||
if body_rules:
|
||||
body = apply_body_rules(body, body_rules)
|
||||
|
||||
# 应用请求头规则(在请求头构建后应用)
|
||||
if header_rules:
|
||||
# 获取认证头名称,防止被规则覆盖
|
||||
from src.core.api_format import get_auth_config_for_endpoint
|
||||
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
|
||||
header_builder = HeaderBuilder()
|
||||
header_builder.add_many(headers)
|
||||
header_builder.apply_rules(header_rules, protected_keys)
|
||||
headers = header_builder.build()
|
||||
|
||||
# 获取有效的模型名称
|
||||
effective_model_name = model_name or request_data.get("model")
|
||||
|
||||
|
||||
@@ -1922,10 +1922,6 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
logger.warning(f"[{ctx.request_id}] 流式请求失败,未选中提供商")
|
||||
return
|
||||
|
||||
# Claude API 的 input_tokens 已经是非缓存部分,不需要再减去 cached_tokens
|
||||
# 实际计费的输入 tokens = input_tokens + cache_creation_tokens(缓存读取免费或折扣)
|
||||
actual_input_tokens = ctx.input_tokens
|
||||
|
||||
# 获取新的 DB session
|
||||
db_gen = get_db()
|
||||
bg_db = next(db_gen)
|
||||
@@ -1987,7 +1983,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
is_stream=True,
|
||||
api_format=ctx.api_format,
|
||||
provider_request_headers=ctx.provider_request_headers,
|
||||
input_tokens=actual_input_tokens,
|
||||
input_tokens=ctx.input_tokens,
|
||||
output_tokens=ctx.output_tokens,
|
||||
cache_creation_tokens=ctx.cache_creation_tokens,
|
||||
cache_read_tokens=ctx.cached_tokens,
|
||||
@@ -2001,7 +1997,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
logger.debug(f"{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:{actual_input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}"
|
||||
f"{ctx.status_code} | in:{ctx.input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}"
|
||||
)
|
||||
else:
|
||||
# 服务端/上游异常:记录为失败
|
||||
@@ -2017,7 +2013,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
api_format=ctx.api_format,
|
||||
provider_request_headers=ctx.provider_request_headers,
|
||||
# 预估 token 信息(来自 message_start 事件)
|
||||
input_tokens=actual_input_tokens,
|
||||
input_tokens=ctx.input_tokens,
|
||||
output_tokens=ctx.output_tokens,
|
||||
cache_creation_tokens=ctx.cache_creation_tokens,
|
||||
cache_read_tokens=ctx.cached_tokens,
|
||||
@@ -2033,7 +2029,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
logger.debug(f"{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:{actual_input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}"
|
||||
f"{ctx.status_code} | in:{ctx.input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}"
|
||||
)
|
||||
else:
|
||||
# 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据
|
||||
@@ -2052,12 +2048,12 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
logger.debug(
|
||||
f"[{ctx.request_id}] 开始记录 Usage: "
|
||||
f"provider={ctx.provider_name}, model={ctx.model}, "
|
||||
f"in={actual_input_tokens}, out={ctx.output_tokens}"
|
||||
f"in={ctx.input_tokens}, out={ctx.output_tokens}"
|
||||
)
|
||||
total_cost = await bg_telemetry.record_success(
|
||||
provider=ctx.provider_name,
|
||||
model=ctx.model,
|
||||
input_tokens=actual_input_tokens,
|
||||
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, # 传递首字时间
|
||||
@@ -2545,8 +2541,6 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
output_tokens = usage.get("output_tokens", 0)
|
||||
cached_tokens = usage.get("cache_read_tokens", 0)
|
||||
cache_creation_tokens = usage.get("cache_creation_tokens", 0)
|
||||
# Claude API 的 input_tokens 已经是非缓存部分,不需要再减去 cached_tokens
|
||||
actual_input_tokens = input_tokens
|
||||
|
||||
output_text = self.parser.extract_text_content(response_json)[:200]
|
||||
|
||||
@@ -2560,7 +2554,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
total_cost = await self.telemetry.record_success(
|
||||
provider=provider_name,
|
||||
model=model,
|
||||
input_tokens=actual_input_tokens,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=status_code,
|
||||
|
||||
@@ -252,6 +252,9 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
api_key: str,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 端点规则参数
|
||||
body_rules: list[dict[str, Any]] | None = None,
|
||||
header_rules: list[dict[str, Any]] | None = None,
|
||||
# 用量计算参数
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
@@ -261,6 +264,10 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
model_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""测试 Gemini API 模型连接性(非流式)"""
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
from src.api.handlers.base.request_builder import apply_body_rules
|
||||
from src.core.api_format.headers import HeaderBuilder
|
||||
|
||||
# Gemini需要从request_data或model_name参数获取model名称
|
||||
effective_model_name = model_name or request_data.get("model", "")
|
||||
if not effective_model_name:
|
||||
@@ -277,8 +284,21 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||
body = cls.build_request_body(request_data)
|
||||
|
||||
# 使用基类的通用endpoint checker
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
|
||||
if body_rules:
|
||||
body = apply_body_rules(body, body_rules)
|
||||
|
||||
# 应用请求头规则(在请求头构建后应用)
|
||||
if header_rules:
|
||||
# 获取认证头名称,防止被规则覆盖
|
||||
from src.core.api_format import get_auth_config_for_endpoint
|
||||
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
|
||||
header_builder = HeaderBuilder()
|
||||
header_builder.add_many(headers)
|
||||
header_builder.apply_rules(header_rules, protected_keys)
|
||||
headers = header_builder.build()
|
||||
|
||||
return await run_endpoint_check(
|
||||
client=client,
|
||||
|
||||
Reference in New Issue
Block a user