Files
Aether/_deprecated_py_src/api/internal/gateway_finalize_chat.py
fawney19 1d9c77522a refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构
- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/
- 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层
- 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构
- 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations)
- 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image
- 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
2026-04-03 16:26:16 +08:00

361 lines
15 KiB
Python

from __future__ import annotations
import base64
import inspect
import time
import uuid
from typing import Any
from fastapi import BackgroundTasks, HTTPException, Request
from fastapi.responses import JSONResponse, Response
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.database import create_session, get_db
from src.models.database import ApiKey, User
from .gateway_contract import GatewaySyncReportRequest
from .gateway_finalize_common import _gateway_module
async def _run_gateway_chat_sync_finalize_background(
payload: GatewaySyncReportRequest,
db: Session,
) -> None:
gateway_module = _gateway_module()
try:
await gateway_module._finalize_gateway_chat_sync(
payload,
db=db,
background_tasks=None,
allow_fast_path=False,
)
except Exception as exc:
logger.warning("gateway background chat finalize failed: {}", exc)
async def _run_gateway_chat_sync_finalize_background_with_session(
payload: GatewaySyncReportRequest,
) -> None:
gateway_module = _gateway_module()
db = create_session()
try:
await gateway_module._run_gateway_chat_sync_finalize_background(payload, db)
finally:
gateway_module._close_gateway_session(db)
async def _finalize_gateway_chat_sync(
payload: GatewaySyncReportRequest,
*,
db: Session,
background_tasks: BackgroundTasks | None = None,
allow_fast_path: bool = True,
) -> Response:
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import is_format_converted
from src.api.handlers.base.utils import (
build_json_response_for_client,
filter_proxy_response_headers,
get_format_converter_registry,
resolve_client_accept_encoding,
)
from src.api.handlers.claude import ClaudeChatAdapter
from src.api.handlers.gemini import GeminiChatAdapter
from src.api.handlers.openai import OpenAIChatAdapter
from src.core.exceptions import EmbeddedErrorException
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.scheduling.schemas import ProviderCandidate
gateway_module = _gateway_module()
context = dict(payload.report_context or {})
user_id = str(context.get("user_id") or "").strip()
api_key_id = str(context.get("api_key_id") or "").strip()
provider_id = str(context.get("provider_id") or "").strip()
endpoint_id = str(context.get("endpoint_id") or "").strip()
key_id = str(context.get("key_id") or "").strip()
request_id = str(context.get("request_id") or payload.trace_id or uuid.uuid4().hex[:8]).strip()
model = str(context.get("model") or "unknown").strip() or "unknown"
provider_api_format = str(context.get("provider_api_format") or "").strip().lower()
client_api_format = str(context.get("client_api_format") or "").strip().lower()
if not all([user_id, api_key_id, provider_id, endpoint_id, key_id, client_api_format]):
raise HTTPException(status_code=400, detail="Missing gateway chat finalize context")
if allow_fast_path and background_tasks is not None:
fast_response = await gateway_module._maybe_build_gateway_core_sync_fast_success_response(
payload
)
if fast_response is not None:
background_tasks.add_task(
gateway_module._run_gateway_chat_sync_finalize_background,
payload.model_copy(deep=True),
db,
)
return fast_response
user = db.query(User).filter(User.id == user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
provider = db.query(Provider).filter(Provider.id == provider_id).first()
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not user or not api_key or not provider or not endpoint or not key:
raise HTTPException(status_code=400, detail="Invalid gateway chat finalize context")
if client_api_format == "claude:chat":
adapter = ClaudeChatAdapter()
elif client_api_format == "gemini:chat":
adapter = GeminiChatAdapter()
else:
adapter = OpenAIChatAdapter()
handler = adapter._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip="127.0.0.1",
user_agent=str(
(context.get("original_headers") or {}).get("user-agent") or "aether-gateway"
),
start_time=time.time(),
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
)
original_headers = dict(context.get("original_headers") or {})
original_request_body = dict(context.get("original_request_body") or {})
mapped_model = str(context.get("mapped_model") or "").strip() or None
candidate = ProviderCandidate(
provider=provider,
endpoint=endpoint,
key=key,
mapping_matched_model=mapped_model,
needs_conversion=is_format_converted(provider_api_format, client_api_format),
provider_api_format=provider_api_format,
)
prep = await handler._prepare_provider_request(
model=model,
provider=provider,
endpoint=endpoint,
key=key,
working_request_body=dict(original_request_body),
original_headers=original_headers,
client_api_format=client_api_format,
provider_api_format=provider_api_format,
candidate=candidate,
client_is_stream=False,
)
response_json: dict[str, Any] | None = None
provider_response_json: dict[str, Any] | None = None
if prep.upstream_is_stream and payload.body_base64 and payload.status_code < 400:
sync_executor = ChatSyncExecutor(handler)
sync_executor._ctx.provider_api_format_for_error = provider_api_format
sync_executor._ctx.client_api_format_for_error = client_api_format
sync_executor._ctx.needs_conversion_for_error = bool(prep.needs_conversion)
sync_executor._ctx.mapped_model_result = mapped_model
try:
response_json = await sync_executor._finalize_rust_stream_sync_result(
prepared_plan=prep,
provider=provider,
model=model,
response_body_bytes=gateway_module._extract_gateway_report_body_bytes(payload),
)
if isinstance(sync_executor._ctx.provider_response_json, dict):
provider_response_json = dict(sync_executor._ctx.provider_response_json)
except EmbeddedErrorException as exc:
payload = payload.model_copy(
update={
"status_code": int(exc.error_code or 400),
"body_json": gateway_module._build_gateway_embedded_error_payload(exc),
"body_base64": None,
}
)
except Exception as exc:
raise HTTPException(
status_code=502,
detail="Invalid upstream chat stream response",
) from exc
provider_error_parser = get_parser_for_format(provider_api_format or client_api_format or "")
is_error_response = payload.status_code >= 400
if isinstance(payload.body_json, dict):
try:
is_error_response = is_error_response or provider_error_parser.is_error_response(
dict(payload.body_json)
)
except Exception:
is_error_response = is_error_response or payload.body_json.get("error") is not None
client_accept_encoding = resolve_client_accept_encoding(original_headers, None)
if is_error_response:
error_status_code = gateway_module._resolve_gateway_sync_error_status_code(
payload,
provider_parser=provider_error_parser,
)
error_payload = gateway_module._build_gateway_sync_error_payload(
payload,
client_api_format=client_api_format,
provider_api_format=provider_api_format or client_api_format,
needs_conversion=bool(prep.needs_conversion),
)
client_response_headers = filter_proxy_response_headers(dict(payload.headers or {}))
client_response_headers["content-type"] = "application/json"
client_response = build_json_response_for_client(
status_code=error_status_code,
content=error_payload,
headers=client_response_headers,
client_accept_encoding=client_accept_encoding,
)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
telemetry_writer = gateway_module._build_gateway_sync_telemetry_writer(
db=db,
request_id=request_id,
user_id=user_id,
api_key_id=api_key_id,
fallback_telemetry=handler.telemetry,
)
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
await gateway_module._schedule_gateway_sync_telemetry(
background_tasks=background_tasks,
telemetry_writer=telemetry_writer,
operation="record_failure",
provider=str(context.get("provider_name") or provider.name or "unknown"),
model=model,
response_time_ms=response_time_ms,
status_code=error_status_code,
error_message=gateway_module._extract_gateway_sync_error_message(payload),
request_headers=original_headers,
request_body=original_request_body,
provider_request_body=context.get("provider_request_body"),
is_stream=False,
api_format=client_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
response_headers=dict(payload.headers or {}),
client_response_headers=dict(client_response.headers),
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, client_api_format),
target_model=mapped_model,
metadata=request_metadata,
)
return client_response
if response_json is None:
if not isinstance(payload.body_json, dict):
raise HTTPException(status_code=502, detail="Invalid upstream chat response")
response_json = dict(payload.body_json)
if prep.envelope:
response_json = prep.envelope.unwrap_response(response_json)
prep.envelope.postprocess_unwrapped_response(model=model, data=response_json)
if prep.needs_conversion:
provider_response_json = dict(response_json)
registry = get_format_converter_registry()
response_json = registry.convert_response(
response_json,
provider_api_format,
client_api_format,
requested_model=model,
)
response_json = handler._normalize_response(response_json)
extract_usage = getattr(handler, "_extract_usage", None)
if callable(extract_usage):
usage_info = extract_usage(response_json)
else:
parser = getattr(handler, "parser", None)
parser_extract_usage = getattr(parser, "extract_usage_from_response", None)
usage_info = parser_extract_usage(response_json) if callable(parser_extract_usage) else {}
extract_response_metadata = getattr(handler, "_extract_response_metadata", None)
response_metadata = (
extract_response_metadata(response_json) if callable(extract_response_metadata) else None
)
client_response_headers = filter_proxy_response_headers(dict(payload.headers or {}))
client_response_headers["content-type"] = "application/json"
client_response = build_json_response_for_client(
status_code=payload.status_code,
content=response_json,
headers=client_response_headers,
client_accept_encoding=client_accept_encoding,
)
request_metadata: dict[str, Any] = {
"gateway_direct_executor": True,
"phase": "3c_trial",
}
proxy_info = context.get("proxy_info")
if isinstance(proxy_info, dict):
request_metadata["proxy"] = proxy_info
telemetry_writer = gateway_module._build_gateway_sync_telemetry_writer(
db=db,
request_id=request_id,
user_id=user_id,
api_key_id=api_key_id,
fallback_telemetry=handler.telemetry,
)
response_time_ms = 0
if isinstance(payload.telemetry, dict):
raw_elapsed_ms = payload.telemetry.get("elapsed_ms")
try:
response_time_ms = max(int(raw_elapsed_ms or 0), 0)
except (TypeError, ValueError):
response_time_ms = 0
await gateway_module._schedule_gateway_sync_telemetry(
background_tasks=background_tasks,
telemetry_writer=telemetry_writer,
operation="record_success",
provider=str(context.get("provider_name") or provider.name or "unknown"),
model=model,
input_tokens=int(usage_info.get("input_tokens", 0) or 0),
output_tokens=int(usage_info.get("output_tokens", 0) or 0),
response_time_ms=response_time_ms,
status_code=payload.status_code,
request_headers=original_headers,
request_body=original_request_body,
response_headers=dict(payload.headers or {}),
client_response_headers=dict(client_response.headers),
response_body=provider_response_json or response_json,
client_response_body=response_json if provider_response_json else None,
provider_request_headers=dict(context.get("provider_request_headers") or {}),
provider_request_body=context.get("provider_request_body"),
is_stream=False,
provider_id=str(context.get("provider_id") or "") or None,
provider_endpoint_id=str(context.get("endpoint_id") or "") or None,
provider_api_key_id=str(context.get("key_id") or "") or None,
api_format=client_api_format,
api_family=adapter.API_FAMILY.value if adapter.API_FAMILY else None,
endpoint_kind=adapter.ENDPOINT_KIND.value if adapter.ENDPOINT_KIND else None,
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, client_api_format),
target_model=mapped_model,
metadata=gateway_module._build_gateway_usage_metadata(
request_metadata=request_metadata,
response_metadata=response_metadata if response_metadata else None,
),
)
return client_response