mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat(usage): 增强 cURL/Replay 对 Vertex AI、OAuth 和 envelope 提供商的支持
- _resolve_provider_auth 返回 decrypted_auth_config 用于 Vertex AI URL 构建 - _build_provider_url_safe 新增 API 格式默认路径回退及模板变量防护 - Replay 适配器基于 api_format 智能匹配端点和 Key - Replay 适配器支持 envelope 包装(kiro/codex/antigravity 等特殊提供商) - 复用 target_provider_obj 减少重复 Provider 查询
This commit is contained in:
@@ -1429,8 +1429,9 @@ def _build_provider_url_safe(
|
|||||||
model_name: str | None,
|
model_name: str | None,
|
||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
provider_key: ProviderAPIKey | None,
|
provider_key: ProviderAPIKey | None,
|
||||||
|
decrypted_auth_config: dict[str, Any] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""构建 Provider URL,build_provider_url 失败时回退到 base_url + custom_path。"""
|
"""构建 Provider URL,build_provider_url 失败时回退到 base_url + custom_path/默认路径。"""
|
||||||
from src.services.provider.transport import build_provider_url
|
from src.services.provider.transport import build_provider_url
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -1439,10 +1440,33 @@ def _build_provider_url_safe(
|
|||||||
path_params={"model": model_name} if model_name else None,
|
path_params={"model": model_name} if model_name else None,
|
||||||
is_stream=is_stream,
|
is_stream=is_stream,
|
||||||
key=provider_key,
|
key=provider_key,
|
||||||
|
decrypted_auth_config=decrypted_auth_config,
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
base = (endpoint.base_url or "").rstrip("/")
|
base = (endpoint.base_url or "").rstrip("/")
|
||||||
return f"{base}{endpoint.custom_path}" if endpoint.custom_path else base
|
if endpoint.custom_path:
|
||||||
|
return f"{base}{endpoint.custom_path}"
|
||||||
|
# 尝试使用 API 格式的默认路径
|
||||||
|
try:
|
||||||
|
from src.core.api_format.metadata import get_default_path_for_endpoint
|
||||||
|
|
||||||
|
ep_sig = (getattr(endpoint, "api_format", "") or "").strip().lower()
|
||||||
|
if ep_sig:
|
||||||
|
path = get_default_path_for_endpoint(ep_sig)
|
||||||
|
if model_name:
|
||||||
|
# Gemini 路径含 {action},回退时固定用非流式操作
|
||||||
|
action = "streamGenerateContent" if is_stream else "generateContent"
|
||||||
|
try:
|
||||||
|
path = path.format(model=model_name, action=action)
|
||||||
|
except KeyError:
|
||||||
|
pass
|
||||||
|
elif "{model}" in path:
|
||||||
|
# model 为空且路径含模板变量,无法构造有效路径,回退到纯 base URL
|
||||||
|
return base
|
||||||
|
return f"{base}{path}"
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return base
|
||||||
|
|
||||||
|
|
||||||
def _build_fresh_headers(
|
def _build_fresh_headers(
|
||||||
@@ -1464,8 +1488,8 @@ async def _resolve_provider_auth(
|
|||||||
provider_key: ProviderAPIKey,
|
provider_key: ProviderAPIKey,
|
||||||
endpoint: ProviderEndpoint,
|
endpoint: ProviderEndpoint,
|
||||||
db: Session,
|
db: Session,
|
||||||
) -> dict[str, str]:
|
) -> tuple[dict[str, str], dict[str, Any] | None]:
|
||||||
"""解析 Provider Key 的认证信息,返回可直接用于请求的认证头字典。
|
"""解析 Provider Key 的认证信息,返回 (认证头字典, 解密后的 auth_config)。
|
||||||
|
|
||||||
支持: api_key / oauth / vertex_ai 三种 auth_type。
|
支持: api_key / oauth / vertex_ai 三种 auth_type。
|
||||||
"""
|
"""
|
||||||
@@ -1474,6 +1498,7 @@ async def _resolve_provider_auth(
|
|||||||
|
|
||||||
auth_type = str(getattr(provider_key, "auth_type", "api_key") or "api_key").lower()
|
auth_type = str(getattr(provider_key, "auth_type", "api_key") or "api_key").lower()
|
||||||
auth_headers: dict[str, str] = {}
|
auth_headers: dict[str, str] = {}
|
||||||
|
decrypted_auth_config: dict[str, Any] | None = None
|
||||||
|
|
||||||
if auth_type == "oauth":
|
if auth_type == "oauth":
|
||||||
from src.services.provider.oauth_token import resolve_oauth_access_token
|
from src.services.provider.oauth_token import resolve_oauth_access_token
|
||||||
@@ -1510,10 +1535,11 @@ async def _resolve_provider_auth(
|
|||||||
)
|
)
|
||||||
access_token = resolved.access_token or ""
|
access_token = resolved.access_token or ""
|
||||||
auth_headers["Authorization"] = f"Bearer {access_token}"
|
auth_headers["Authorization"] = f"Bearer {access_token}"
|
||||||
|
decrypted_auth_config = resolved.decrypted_auth_config
|
||||||
|
|
||||||
# Codex 等需要 account_id
|
# Codex 等需要 account_id
|
||||||
if resolved.decrypted_auth_config:
|
if decrypted_auth_config:
|
||||||
account_id = resolved.decrypted_auth_config.get("account_id")
|
account_id = decrypted_auth_config.get("account_id")
|
||||||
if account_id:
|
if account_id:
|
||||||
auth_headers["chatgpt-account-id"] = str(account_id)
|
auth_headers["chatgpt-account-id"] = str(account_id)
|
||||||
|
|
||||||
@@ -1544,7 +1570,7 @@ async def _resolve_provider_auth(
|
|||||||
auth_value = f"Bearer {decrypted_key}" if auth_type_cfg == "bearer" else decrypted_key
|
auth_value = f"Bearer {decrypted_key}" if auth_type_cfg == "bearer" else decrypted_key
|
||||||
auth_headers[auth_header] = auth_value
|
auth_headers[auth_header] = auth_value
|
||||||
|
|
||||||
return auth_headers
|
return auth_headers, decrypted_auth_config
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -1588,21 +1614,17 @@ class AdminUsageCurlAdapter(AdminApiAdapter):
|
|||||||
if key_id:
|
if key_id:
|
||||||
provider_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
provider_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
||||||
|
|
||||||
# 重建请求 URL
|
# 解析认证信息
|
||||||
url: str | None = None
|
|
||||||
if endpoint:
|
|
||||||
model_name = usage_record.target_model or usage_record.model
|
|
||||||
url = _build_provider_url_safe(
|
|
||||||
endpoint, model_name, usage_record.is_stream or False, provider_key
|
|
||||||
)
|
|
||||||
|
|
||||||
# 解析认证信息并构建请求头
|
|
||||||
stored_headers = usage_record.provider_request_headers or {}
|
stored_headers = usage_record.provider_request_headers or {}
|
||||||
headers: dict[str, str] = {}
|
headers: dict[str, str] = {}
|
||||||
|
auth_headers: dict[str, str] = {}
|
||||||
|
decrypted_auth_config: dict[str, Any] | None = None
|
||||||
|
|
||||||
if provider_key and endpoint:
|
if provider_key and endpoint:
|
||||||
try:
|
try:
|
||||||
auth_headers = await _resolve_provider_auth(provider_key, endpoint, db)
|
auth_headers, decrypted_auth_config = await _resolve_provider_auth(
|
||||||
|
provider_key, endpoint, db
|
||||||
|
)
|
||||||
|
|
||||||
if stored_headers:
|
if stored_headers:
|
||||||
# 有存储的请求头:替换被脱敏的认证头为真实值
|
# 有存储的请求头:替换被脱敏的认证头为真实值
|
||||||
@@ -1619,6 +1641,18 @@ class AdminUsageCurlAdapter(AdminApiAdapter):
|
|||||||
else:
|
else:
|
||||||
headers = dict(stored_headers)
|
headers = dict(stored_headers)
|
||||||
|
|
||||||
|
# 重建请求 URL(在认证解析之后,以便传递 decrypted_auth_config 给 Vertex AI 等场景)
|
||||||
|
url: str | None = None
|
||||||
|
if endpoint:
|
||||||
|
model_name = usage_record.target_model or usage_record.model
|
||||||
|
url = _build_provider_url_safe(
|
||||||
|
endpoint,
|
||||||
|
model_name,
|
||||||
|
usage_record.is_stream or False,
|
||||||
|
provider_key,
|
||||||
|
decrypted_auth_config,
|
||||||
|
)
|
||||||
|
|
||||||
# 确保始终有 Content-Type
|
# 确保始终有 Content-Type
|
||||||
if not any(k.lower() == "content-type" for k in headers):
|
if not any(k.lower() == "content-type" for k in headers):
|
||||||
headers["Content-Type"] = "application/json"
|
headers["Content-Type"] = "application/json"
|
||||||
@@ -1673,8 +1707,17 @@ class AdminUsageReplayAdapter(AdminApiAdapter):
|
|||||||
db = context.db
|
db = context.db
|
||||||
usage_record = _find_usage_record(db, self.usage_id)
|
usage_record = _find_usage_record(db, self.usage_id)
|
||||||
|
|
||||||
|
# 确定原始请求的 API 格式(用于端点/Key 匹配)
|
||||||
|
original_api_format = (
|
||||||
|
(usage_record.endpoint_api_format or usage_record.api_format or "").strip().lower()
|
||||||
|
)
|
||||||
|
original_api_family = (
|
||||||
|
original_api_format.split(":")[0] if ":" in original_api_format else ""
|
||||||
|
)
|
||||||
|
|
||||||
# 确定目标端点和密钥
|
# 确定目标端点和密钥
|
||||||
target_pid = self.target_provider_id
|
target_pid = self.target_provider_id
|
||||||
|
target_provider_obj: Provider | None = None
|
||||||
if self.target_endpoint_id:
|
if self.target_endpoint_id:
|
||||||
endpoint = (
|
endpoint = (
|
||||||
db.query(ProviderEndpoint)
|
db.query(ProviderEndpoint)
|
||||||
@@ -1685,17 +1728,35 @@ class AdminUsageReplayAdapter(AdminApiAdapter):
|
|||||||
raise HTTPException(status_code=404, detail="Target endpoint not found")
|
raise HTTPException(status_code=404, detail="Target endpoint not found")
|
||||||
target_pid = str(endpoint.provider_id)
|
target_pid = str(endpoint.provider_id)
|
||||||
elif target_pid:
|
elif target_pid:
|
||||||
target_provider = db.query(Provider).filter(Provider.id == target_pid).first()
|
target_provider_obj = db.query(Provider).filter(Provider.id == target_pid).first()
|
||||||
if not target_provider:
|
if not target_provider_obj:
|
||||||
raise HTTPException(status_code=404, detail="Target provider not found")
|
raise HTTPException(status_code=404, detail="Target provider not found")
|
||||||
endpoint = (
|
# 优先匹配相同 api_format 的端点,其次匹配同 family,最后取任意 active 端点
|
||||||
|
active_endpoints = (
|
||||||
db.query(ProviderEndpoint)
|
db.query(ProviderEndpoint)
|
||||||
.filter(
|
.filter(
|
||||||
ProviderEndpoint.provider_id == target_pid,
|
ProviderEndpoint.provider_id == target_pid,
|
||||||
ProviderEndpoint.is_active == True, # noqa: E712
|
ProviderEndpoint.is_active == True, # noqa: E712
|
||||||
)
|
)
|
||||||
.first()
|
.all()
|
||||||
)
|
)
|
||||||
|
endpoint = None
|
||||||
|
if active_endpoints and original_api_format:
|
||||||
|
# 精确匹配 api_format
|
||||||
|
for ep in active_endpoints:
|
||||||
|
ep_fmt = (getattr(ep, "api_format", "") or "").strip().lower()
|
||||||
|
if ep_fmt == original_api_format:
|
||||||
|
endpoint = ep
|
||||||
|
break
|
||||||
|
# 同 family 匹配
|
||||||
|
if not endpoint and original_api_family:
|
||||||
|
for ep in active_endpoints:
|
||||||
|
ep_family = (getattr(ep, "api_family", "") or "").strip().lower()
|
||||||
|
if ep_family == original_api_family:
|
||||||
|
endpoint = ep
|
||||||
|
break
|
||||||
|
if not endpoint and active_endpoints:
|
||||||
|
endpoint = active_endpoints[0]
|
||||||
if not endpoint:
|
if not endpoint:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=404, detail="No active endpoint found for target provider"
|
status_code=404, detail="No active endpoint found for target provider"
|
||||||
@@ -1715,20 +1776,36 @@ class AdminUsageReplayAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 确定 API Key
|
# 确定 API Key
|
||||||
|
target_ep_format = (getattr(endpoint, "api_format", "") or "").strip().lower()
|
||||||
provider_key = None
|
provider_key = None
|
||||||
if self.target_api_key_id:
|
if self.target_api_key_id:
|
||||||
provider_key = (
|
provider_key = (
|
||||||
db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.target_api_key_id).first()
|
db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.target_api_key_id).first()
|
||||||
)
|
)
|
||||||
elif target_pid:
|
elif target_pid:
|
||||||
provider_key = (
|
# 优先选择 api_formats 包含目标端点格式的 Key
|
||||||
|
active_keys = (
|
||||||
db.query(ProviderAPIKey)
|
db.query(ProviderAPIKey)
|
||||||
.filter(
|
.filter(
|
||||||
ProviderAPIKey.provider_id == target_pid,
|
ProviderAPIKey.provider_id == target_pid,
|
||||||
ProviderAPIKey.is_active == True, # noqa: E712
|
ProviderAPIKey.is_active == True, # noqa: E712
|
||||||
)
|
)
|
||||||
.first()
|
.all()
|
||||||
)
|
)
|
||||||
|
if active_keys and target_ep_format:
|
||||||
|
for k in active_keys:
|
||||||
|
k_formats = getattr(k, "api_formats", None)
|
||||||
|
if k_formats is None:
|
||||||
|
# None 表示支持所有格式,直接选中
|
||||||
|
provider_key = k
|
||||||
|
break
|
||||||
|
if isinstance(k_formats, list) and target_ep_format in [
|
||||||
|
f.strip().lower() for f in k_formats
|
||||||
|
]:
|
||||||
|
provider_key = k
|
||||||
|
break
|
||||||
|
if not provider_key and active_keys:
|
||||||
|
provider_key = active_keys[0]
|
||||||
else:
|
else:
|
||||||
if usage_record.provider_api_key_id:
|
if usage_record.provider_api_key_id:
|
||||||
provider_key = (
|
provider_key = (
|
||||||
@@ -1742,14 +1819,18 @@ class AdminUsageReplayAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
# 根据 auth_type 正确解析认证(支持 OAuth / Vertex AI / API Key)
|
# 根据 auth_type 正确解析认证(支持 OAuth / Vertex AI / API Key)
|
||||||
try:
|
try:
|
||||||
auth_headers = await _resolve_provider_auth(provider_key, endpoint, db)
|
auth_headers, decrypted_auth_config = await _resolve_provider_auth(
|
||||||
|
provider_key, endpoint, db
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("[replay] Failed to resolve auth for key {}: {}", provider_key.id, e)
|
logger.error("[replay] Failed to resolve auth for key {}: {}", provider_key.id, e)
|
||||||
raise HTTPException(status_code=500, detail="Failed to resolve provider authentication")
|
raise HTTPException(status_code=500, detail="Failed to resolve provider authentication")
|
||||||
|
|
||||||
# 构建 URL
|
# 构建 URL
|
||||||
model_name = usage_record.target_model or usage_record.model
|
model_name = usage_record.target_model or usage_record.model
|
||||||
url = _build_provider_url_safe(endpoint, model_name, False, provider_key)
|
url = _build_provider_url_safe(
|
||||||
|
endpoint, model_name, False, provider_key, decrypted_auth_config
|
||||||
|
)
|
||||||
|
|
||||||
# 构建请求头(Content-Type + 认证头 + 端点额外头)
|
# 构建请求头(Content-Type + 认证头 + 端点额外头)
|
||||||
headers = _build_fresh_headers(auth_headers, endpoint)
|
headers = _build_fresh_headers(auth_headers, endpoint)
|
||||||
@@ -1759,25 +1840,21 @@ class AdminUsageReplayAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
# 格式转换:如果存储的请求体格式与目标端点格式不同,需要转换
|
# 格式转换:如果存储的请求体格式与目标端点格式不同,需要转换
|
||||||
if isinstance(body, dict):
|
if isinstance(body, dict):
|
||||||
# 确定存储体的格式(发送给原 Provider 的格式)
|
|
||||||
stored_format = (
|
|
||||||
(usage_record.endpoint_api_format or usage_record.api_format or "").strip().lower()
|
|
||||||
)
|
|
||||||
target_format = str(getattr(endpoint, "api_format", "") or "").strip().lower()
|
target_format = str(getattr(endpoint, "api_format", "") or "").strip().lower()
|
||||||
|
|
||||||
if stored_format and target_format and stored_format != target_format:
|
if original_api_format and target_format and original_api_format != target_format:
|
||||||
try:
|
try:
|
||||||
from src.core.api_format.conversion import format_conversion_registry
|
from src.core.api_format.conversion import format_conversion_registry
|
||||||
|
|
||||||
body = format_conversion_registry.convert_request(
|
body = format_conversion_registry.convert_request(
|
||||||
body,
|
body,
|
||||||
source_format=stored_format,
|
source_format=original_api_format,
|
||||||
target_format=target_format,
|
target_format=target_format,
|
||||||
)
|
)
|
||||||
except Exception as conv_err:
|
except Exception as conv_err:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"[replay] Format conversion {} -> {} failed: {}",
|
"[replay] Format conversion {} -> {} failed: {}",
|
||||||
stored_format,
|
original_api_format,
|
||||||
target_format,
|
target_format,
|
||||||
conv_err,
|
conv_err,
|
||||||
)
|
)
|
||||||
@@ -1793,16 +1870,45 @@ class AdminUsageReplayAdapter(AdminApiAdapter):
|
|||||||
else:
|
else:
|
||||||
body["stream"] = False
|
body["stream"] = False
|
||||||
|
|
||||||
|
# 反代提供商 envelope 包装:kiro/codex/antigravity 等需要特殊的请求体格式和额外请求头
|
||||||
|
if not target_provider_obj:
|
||||||
|
target_provider_obj = (
|
||||||
|
db.query(Provider).filter(Provider.id == endpoint.provider_id).first()
|
||||||
|
)
|
||||||
|
if isinstance(body, dict):
|
||||||
|
try:
|
||||||
|
from src.services.provider.envelope import get_provider_envelope
|
||||||
|
|
||||||
|
target_provider_type = (
|
||||||
|
str(getattr(target_provider_obj, "provider_type", "") or "").lower()
|
||||||
|
if target_provider_obj
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
target_ep_sig = (getattr(endpoint, "api_format", "") or "").strip().lower()
|
||||||
|
|
||||||
|
envelope = get_provider_envelope(
|
||||||
|
provider_type=target_provider_type,
|
||||||
|
endpoint_sig=target_ep_sig,
|
||||||
|
)
|
||||||
|
if envelope:
|
||||||
|
body, _ = envelope.wrap_request(
|
||||||
|
body,
|
||||||
|
model=model_name or "",
|
||||||
|
url_model=model_name,
|
||||||
|
decrypted_auth_config=decrypted_auth_config,
|
||||||
|
)
|
||||||
|
# envelope 可能注入额外请求头(如 Kiro 的 AWS 签名头、Codex 的 OAuth 头)
|
||||||
|
extra_envelope_headers = envelope.extra_headers()
|
||||||
|
if extra_envelope_headers:
|
||||||
|
headers.update(extra_envelope_headers)
|
||||||
|
except Exception as env_err:
|
||||||
|
logger.warning("[replay] Envelope wrap failed: {}", env_err)
|
||||||
|
# envelope 失败仍发送原始体
|
||||||
|
|
||||||
# 获取提供商名称
|
# 获取提供商名称
|
||||||
provider_name = usage_record.provider_name
|
provider_name = usage_record.provider_name
|
||||||
if self.target_provider_id:
|
if target_provider_obj:
|
||||||
target_p = db.query(Provider).filter(Provider.id == self.target_provider_id).first()
|
provider_name = target_provider_obj.name
|
||||||
if target_p:
|
|
||||||
provider_name = target_p.name
|
|
||||||
elif endpoint and endpoint.provider_id:
|
|
||||||
p = db.query(Provider).filter(Provider.id == endpoint.provider_id).first()
|
|
||||||
if p:
|
|
||||||
provider_name = p.name
|
|
||||||
|
|
||||||
# 发送请求
|
# 发送请求
|
||||||
try:
|
try:
|
||||||
@@ -1812,13 +1918,8 @@ class AdminUsageReplayAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 解析代理(key > provider > 系统默认)
|
# 解析代理(key > provider > 系统默认)
|
||||||
replay_provider = (
|
|
||||||
db.query(Provider).filter(Provider.id == endpoint.provider_id).first()
|
|
||||||
if endpoint
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
eff_proxy = resolve_effective_proxy(
|
eff_proxy = resolve_effective_proxy(
|
||||||
getattr(replay_provider, "proxy", None) if replay_provider else None,
|
getattr(target_provider_obj, "proxy", None) if target_provider_obj else None,
|
||||||
getattr(provider_key, "proxy", None) if provider_key else None,
|
getattr(provider_key, "proxy", None) if provider_key else None,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user