mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
perf: 同步 DB 操作迁移到 asyncio.to_thread,避免阻塞事件循环
failover/stream_telemetry/recording/stream 中的同步 DB 操作(commit/execute/query) 会阻塞 asyncio 事件循环,导致 Hub PING 心跳无法发送、worker idle timeout 断连。 将这些操作包装到 asyncio.to_thread() 中执行。 同时提取 Alembic 迁移脚本中重复的幂等性辅助函数到 alembic/helpers.py, 用批量查询缓存替代逐条 information_schema 查询,backfill SQL 合并为 LEFT JOIN。 新增 failover 中客户端断连的快速终止路径,避免继续无意义的重试。
This commit is contained in:
@@ -24,6 +24,17 @@ from .policy import FailoverAction, RetryMode, RetryPolicy, SkipPolicy
|
||||
from .recorder import CandidateRecorder
|
||||
from .schema import CandidateKey
|
||||
|
||||
_DISCONNECT_EXCEPTION_NAMES = frozenset({"ClientDisconnectedException"})
|
||||
|
||||
|
||||
def _is_client_disconnected(exc: Exception) -> bool:
|
||||
"""检查异常(或其 cause)是否为客户端断连,避免循环 import。"""
|
||||
for obj in (exc, getattr(exc, "cause", None)):
|
||||
if obj is not None and type(obj).__name__ in _DISCONNECT_EXCEPTION_NAMES:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
_SENSITIVE_PATTERN = re.compile(
|
||||
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
|
||||
re.IGNORECASE,
|
||||
@@ -66,6 +77,14 @@ class FailoverEngine:
|
||||
self._error_classifier = error_classifier or ErrorClassifier(db=db)
|
||||
self._recorder = recorder or CandidateRecorder(db)
|
||||
|
||||
async def _db_op(self, func: Callable[[], Any]) -> Any:
|
||||
"""将同步 DB 操作放到线程池执行,避免阻塞 asyncio 事件循环。
|
||||
|
||||
当事件循环被同步 db.commit() / db.execute() 阻塞时,
|
||||
Hub PING 心跳无法发送,导致 worker idle timeout 断连,整个服务不可用。
|
||||
"""
|
||||
return await asyncio.to_thread(func)
|
||||
|
||||
@staticmethod
|
||||
def _collect_error_messages(error: Exception | None) -> str:
|
||||
if error is None:
|
||||
@@ -192,7 +211,7 @@ class FailoverEngine:
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def _mark_remaining_cancelled(
|
||||
async def _mark_remaining_cancelled(
|
||||
self,
|
||||
*,
|
||||
candidate_record_map: dict[tuple[int, int], str] | None,
|
||||
@@ -204,8 +223,8 @@ class FailoverEngine:
|
||||
if not candidate_record_map:
|
||||
return
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
updated = False
|
||||
# Pre-compute outside thread (reads ProviderCandidate attrs that may not be thread-safe)
|
||||
record_ids: list[str] = []
|
||||
for candidate_idx, cand in enumerate(candidates):
|
||||
if candidate_idx < from_candidate_idx:
|
||||
continue
|
||||
@@ -214,11 +233,18 @@ class FailoverEngine:
|
||||
if candidate_idx == from_candidate_idx and retry_idx < from_retry_idx:
|
||||
continue
|
||||
record_id = candidate_record_map.get((candidate_idx, retry_idx))
|
||||
if not record_id:
|
||||
continue
|
||||
if record_id:
|
||||
record_ids.append(record_id)
|
||||
|
||||
if not record_ids:
|
||||
return
|
||||
|
||||
def _do() -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
for rid in record_ids:
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id == record_id)
|
||||
.where(RequestCandidate.id == rid)
|
||||
.where(RequestCandidate.status.in_(["available", "pending"]))
|
||||
.values(
|
||||
status="cancelled",
|
||||
@@ -227,11 +253,10 @@ class FailoverEngine:
|
||||
finished_at=now,
|
||||
)
|
||||
)
|
||||
updated = True
|
||||
|
||||
if updated:
|
||||
self.db.commit()
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
def _append_cancelled_fallback_candidate_keys(
|
||||
self,
|
||||
*,
|
||||
@@ -295,7 +320,7 @@ class FailoverEngine:
|
||||
from_candidate_idx,
|
||||
from_retry_idx,
|
||||
)
|
||||
self._mark_remaining_cancelled(
|
||||
await self._mark_remaining_cancelled(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidates=candidates,
|
||||
from_candidate_idx=from_candidate_idx,
|
||||
@@ -322,6 +347,46 @@ class FailoverEngine:
|
||||
attempt_count=attempt_count,
|
||||
)
|
||||
|
||||
async def _build_disconnected_result(
|
||||
self,
|
||||
*,
|
||||
candidate_record_map: dict[tuple[int, int], str] | None,
|
||||
candidate_keys_fallback: list[CandidateKey],
|
||||
candidates: list[ProviderCandidate],
|
||||
from_candidate_idx: int,
|
||||
from_retry_idx: int,
|
||||
retry_policy: RetryPolicy,
|
||||
request_id: str | None,
|
||||
attempt_count: int,
|
||||
) -> ExecutionResult:
|
||||
"""客户端已断连,立即终止故障转移并返回结果。"""
|
||||
await self._mark_remaining_cancelled(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidates=candidates,
|
||||
from_candidate_idx=from_candidate_idx,
|
||||
from_retry_idx=from_retry_idx,
|
||||
retry_policy=retry_policy,
|
||||
)
|
||||
self._append_cancelled_fallback_candidate_keys(
|
||||
fallback=candidate_keys_fallback,
|
||||
candidates=candidates,
|
||||
from_candidate_idx=from_candidate_idx,
|
||||
from_retry_idx=from_retry_idx,
|
||||
retry_policy=retry_policy,
|
||||
)
|
||||
return ExecutionResult(
|
||||
success=False,
|
||||
error_type="ClientDisconnected",
|
||||
error_message="client_disconnected",
|
||||
last_status_code=499,
|
||||
candidate_keys=self._get_candidate_keys(
|
||||
request_id=request_id,
|
||||
fallback=candidate_keys_fallback,
|
||||
candidates=candidates,
|
||||
),
|
||||
attempt_count=attempt_count,
|
||||
)
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
*,
|
||||
@@ -391,7 +456,7 @@ class FailoverEngine:
|
||||
if should_skip:
|
||||
# PRE_EXPAND: mark all retry slots skipped.
|
||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||
self._mark_candidate_skipped(
|
||||
await self._mark_candidate_skipped(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidate_index=candidate_index,
|
||||
candidate=candidate,
|
||||
@@ -490,17 +555,9 @@ class FailoverEngine:
|
||||
candidate, candidate_index, retry_index, record_id, attempt_count, max_attempts
|
||||
)
|
||||
|
||||
# Mark pending
|
||||
now = datetime.now(timezone.utc)
|
||||
if record_id:
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="pending",
|
||||
started_at=now,
|
||||
)
|
||||
|
||||
# Commit BEFORE await (avoid holding DB connections during slow upstream calls)
|
||||
self._commit_before_await()
|
||||
# Mark pending + commit BEFORE await
|
||||
# (avoid holding DB connections during slow upstream calls)
|
||||
await self._mark_pending_and_commit(record_id)
|
||||
|
||||
try:
|
||||
attempt_result = await self._execute_attempt(
|
||||
@@ -512,7 +569,7 @@ class FailoverEngine:
|
||||
|
||||
# PRE_EXPAND: mark unused slots after request ends (success)
|
||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||
self._mark_remaining_slots_unused(
|
||||
await self._mark_remaining_slots_unused(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidates=candidates,
|
||||
success_candidate_idx=candidate_index,
|
||||
@@ -542,7 +599,7 @@ class FailoverEngine:
|
||||
|
||||
except StreamProbeError as exc:
|
||||
last_status_code = exc.http_status
|
||||
self._record_attempt_failure(record_id, exc, exc.http_status)
|
||||
await self._record_attempt_failure(record_id, exc, exc.http_status)
|
||||
action = FailoverAction.CONTINUE
|
||||
consecutive_failures += 1
|
||||
await self._apply_retry_pacing(
|
||||
@@ -553,6 +610,19 @@ class FailoverEngine:
|
||||
)
|
||||
|
||||
except Exception as exc:
|
||||
if _is_client_disconnected(exc):
|
||||
await self._record_attempt_failure(record_id, exc, 499)
|
||||
return await self._build_disconnected_result(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidate_keys_fallback=candidate_keys_fallback,
|
||||
candidates=candidates,
|
||||
from_candidate_idx=candidate_index,
|
||||
from_retry_idx=retry_index + 1,
|
||||
retry_policy=retry_policy,
|
||||
request_id=request_id,
|
||||
attempt_count=attempt_count,
|
||||
)
|
||||
|
||||
outcome = await self._handle_attempt_error(
|
||||
exc,
|
||||
candidate=candidate,
|
||||
@@ -587,7 +657,7 @@ class FailoverEngine:
|
||||
if action == FailoverAction.CONTINUE:
|
||||
# PRE_EXPAND: if we break early, mark remaining retries of this candidate unused.
|
||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||
self._mark_candidate_remaining_retries_unused(
|
||||
await self._mark_candidate_remaining_retries_unused(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidate_idx=candidate_index,
|
||||
from_retry_idx=retry_index + 1,
|
||||
@@ -603,7 +673,7 @@ class FailoverEngine:
|
||||
|
||||
# exhausted: PRE_EXPAND should not leave 'available' records behind
|
||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||
self._mark_all_remaining_available_unused(candidate_record_map)
|
||||
await self._mark_all_remaining_available_unused(candidate_record_map)
|
||||
|
||||
return ExecutionResult(
|
||||
success=False,
|
||||
@@ -670,7 +740,7 @@ class FailoverEngine:
|
||||
or "pool_skipped"
|
||||
)
|
||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||
self._mark_retry_indices_status(
|
||||
await self._mark_retry_indices_status(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidate_idx=candidate_index,
|
||||
retry_indices=range(
|
||||
@@ -749,15 +819,7 @@ class FailoverEngine:
|
||||
max_attempts,
|
||||
)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
if record_id:
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="pending",
|
||||
started_at=now,
|
||||
)
|
||||
|
||||
self._commit_before_await()
|
||||
await self._mark_pending_and_commit(record_id)
|
||||
|
||||
try:
|
||||
attempt_result = await self._execute_attempt(
|
||||
@@ -768,7 +830,7 @@ class FailoverEngine:
|
||||
last_status_code = int(getattr(attempt_result, "http_status", 0) or 0)
|
||||
|
||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||
self._mark_remaining_slots_unused(
|
||||
await self._mark_remaining_slots_unused(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidates=candidates,
|
||||
success_candidate_idx=candidate_index,
|
||||
@@ -803,7 +865,7 @@ class FailoverEngine:
|
||||
|
||||
except StreamProbeError as exc:
|
||||
last_status_code = exc.http_status
|
||||
self._record_attempt_failure(record_id, exc, exc.http_status)
|
||||
await self._record_attempt_failure(record_id, exc, exc.http_status)
|
||||
action = FailoverAction.CONTINUE
|
||||
consecutive_failures += 1
|
||||
await self._apply_retry_pacing(
|
||||
@@ -814,6 +876,24 @@ class FailoverEngine:
|
||||
)
|
||||
|
||||
except Exception as exc:
|
||||
if _is_client_disconnected(exc):
|
||||
await self._record_attempt_failure(record_id, exc, 499)
|
||||
return (
|
||||
await self._build_disconnected_result(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidate_keys_fallback=candidate_keys_fallback,
|
||||
candidates=candidates,
|
||||
from_candidate_idx=candidate_index,
|
||||
from_retry_idx=composite_retry_index + 1,
|
||||
retry_policy=retry_policy,
|
||||
request_id=request_id,
|
||||
attempt_count=attempt_count,
|
||||
),
|
||||
attempt_count,
|
||||
consecutive_failures,
|
||||
499,
|
||||
)
|
||||
|
||||
outcome = await self._handle_pool_attempt_error(
|
||||
exc,
|
||||
candidate=candidate,
|
||||
@@ -843,7 +923,7 @@ class FailoverEngine:
|
||||
# max_retries_for_key may have been shrunk by error handler;
|
||||
# mark unused up to the *original* retry_slots_per_key to cover
|
||||
# all pre-created records.
|
||||
self._mark_retry_indices_status(
|
||||
await self._mark_retry_indices_status(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidate_idx=candidate_index,
|
||||
retry_indices=range(
|
||||
@@ -862,7 +942,7 @@ class FailoverEngine:
|
||||
# should not terminate the entire request because other providers may still succeed.
|
||||
# When handler_used=True, TaskService raises directly for true STOP semantics.
|
||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||
self._mark_candidate_remaining_retries_unused(
|
||||
await self._mark_candidate_remaining_retries_unused(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidate_idx=candidate_index,
|
||||
from_retry_idx=composite_retry_index + 1,
|
||||
@@ -959,7 +1039,7 @@ class FailoverEngine:
|
||||
candidate, is_success=True, response_text=body_text
|
||||
)
|
||||
if rule_action == FailoverAction.CONTINUE:
|
||||
self._record_attempt_failure(
|
||||
await self._record_attempt_failure(
|
||||
record_id,
|
||||
Exception("success_failover_pattern matched"),
|
||||
200,
|
||||
@@ -969,43 +1049,55 @@ class FailoverEngine:
|
||||
http_status=200,
|
||||
)
|
||||
|
||||
self._record_attempt_success(record_id, attempt_result)
|
||||
await self._record_attempt_success(record_id, attempt_result)
|
||||
return attempt_result
|
||||
|
||||
def _record_attempt_success(self, record_id: str | None, attempt_result: AttemptResult) -> None:
|
||||
async def _record_attempt_success(
|
||||
self, record_id: str | None, attempt_result: AttemptResult
|
||||
) -> None:
|
||||
"""Mark attempt record as success/streaming."""
|
||||
if not record_id:
|
||||
return
|
||||
if attempt_result.kind == AttemptKind.STREAM:
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="streaming",
|
||||
status_code=attempt_result.http_status,
|
||||
)
|
||||
else:
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="success",
|
||||
status_code=attempt_result.http_status,
|
||||
finished_at=datetime.now(timezone.utc),
|
||||
)
|
||||
self.db.commit()
|
||||
|
||||
def _record_attempt_failure(
|
||||
def _do() -> None:
|
||||
if attempt_result.kind == AttemptKind.STREAM:
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="streaming",
|
||||
status_code=attempt_result.http_status,
|
||||
)
|
||||
else:
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="success",
|
||||
status_code=attempt_result.http_status,
|
||||
finished_at=datetime.now(timezone.utc),
|
||||
)
|
||||
self.db.commit()
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
async def _record_attempt_failure(
|
||||
self, record_id: str | None, exc: Exception, status_code: int | None = None
|
||||
) -> None:
|
||||
"""Mark attempt record as failed."""
|
||||
if not record_id:
|
||||
return
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="failed",
|
||||
status_code=status_code,
|
||||
error_type=type(exc).__name__,
|
||||
error_message=self._sanitize(str(exc)),
|
||||
finished_at=datetime.now(timezone.utc),
|
||||
)
|
||||
self.db.commit()
|
||||
error_type = type(exc).__name__
|
||||
error_message = self._sanitize(str(exc))
|
||||
|
||||
def _do() -> None:
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="failed",
|
||||
status_code=status_code,
|
||||
error_type=error_type,
|
||||
error_message=error_message,
|
||||
finished_at=datetime.now(timezone.utc),
|
||||
)
|
||||
self.db.commit()
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
async def _handle_attempt_error(
|
||||
self,
|
||||
@@ -1046,7 +1138,7 @@ class FailoverEngine:
|
||||
|
||||
if outcome.action == FailoverAction.STOP:
|
||||
if retry_policy.mode == RetryMode.PRE_EXPAND and candidate_record_map:
|
||||
self._mark_remaining_slots_unused(
|
||||
await self._mark_remaining_slots_unused(
|
||||
candidate_record_map=candidate_record_map,
|
||||
candidates=candidates,
|
||||
success_candidate_idx=candidate_index,
|
||||
@@ -1120,7 +1212,7 @@ class FailoverEngine:
|
||||
last_status_code = int(getattr(exc, "status_code", 0) or 0) or int(
|
||||
getattr(exc, "http_status", 0) or 0
|
||||
)
|
||||
self._record_attempt_failure(record_id, exc, last_status_code or None)
|
||||
await self._record_attempt_failure(record_id, exc, last_status_code or None)
|
||||
|
||||
return AttemptErrorOutcome(
|
||||
action=action,
|
||||
@@ -1194,13 +1286,33 @@ class FailoverEngine:
|
||||
)
|
||||
return result
|
||||
|
||||
def _commit_before_await(self) -> None:
|
||||
if self.db.in_transaction():
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise
|
||||
async def _commit_before_await(self) -> None:
|
||||
def _do() -> None:
|
||||
if self.db.in_transaction():
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
async def _mark_pending_and_commit(self, record_id: str | None) -> None:
|
||||
"""Mark record as pending and commit, all within a worker thread."""
|
||||
|
||||
def _do() -> None:
|
||||
if record_id:
|
||||
self._update_record(
|
||||
record_id, status="pending", started_at=datetime.now(timezone.utc)
|
||||
)
|
||||
if self.db.in_transaction():
|
||||
try:
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
def _update_record(self, record_id: str, /, **values: Any) -> None:
|
||||
self.db.execute(
|
||||
@@ -1221,23 +1333,31 @@ class FailoverEngine:
|
||||
) -> str:
|
||||
# Create "available" record, then caller will mark pending.
|
||||
extra = self._build_pool_extra_data(candidate)
|
||||
row = RequestCandidateService.create_candidate(
|
||||
db=self.db,
|
||||
request_id=request_id,
|
||||
candidate_index=candidate_index,
|
||||
retry_index=retry_index,
|
||||
user_id=user_id,
|
||||
api_key_id=api_key_id,
|
||||
username=username,
|
||||
api_key_name=api_key_name,
|
||||
provider_id=str(candidate.provider.id),
|
||||
endpoint_id=str(candidate.endpoint.id),
|
||||
key_id=str(candidate.key.id),
|
||||
status="available",
|
||||
is_cached=bool(getattr(candidate, "is_cached", False)),
|
||||
extra_data=extra,
|
||||
)
|
||||
return str(row.id)
|
||||
provider_id = str(candidate.provider.id)
|
||||
endpoint_id = str(candidate.endpoint.id)
|
||||
key_id = str(candidate.key.id)
|
||||
is_cached = bool(getattr(candidate, "is_cached", False))
|
||||
|
||||
def _do() -> str:
|
||||
row = RequestCandidateService.create_candidate(
|
||||
db=self.db,
|
||||
request_id=request_id,
|
||||
candidate_index=candidate_index,
|
||||
retry_index=retry_index,
|
||||
user_id=user_id,
|
||||
api_key_id=api_key_id,
|
||||
username=username,
|
||||
api_key_name=api_key_name,
|
||||
provider_id=provider_id,
|
||||
endpoint_id=endpoint_id,
|
||||
key_id=key_id,
|
||||
status="available",
|
||||
is_cached=is_cached,
|
||||
extra_data=extra,
|
||||
)
|
||||
return str(row.id)
|
||||
|
||||
return await self._db_op(_do)
|
||||
|
||||
async def _create_skipped_record(
|
||||
self,
|
||||
@@ -1253,27 +1373,35 @@ class FailoverEngine:
|
||||
skip_reason: str | None,
|
||||
) -> str:
|
||||
extra = self._build_pool_extra_data(candidate)
|
||||
row = RequestCandidateService.create_candidate(
|
||||
db=self.db,
|
||||
request_id=request_id,
|
||||
candidate_index=candidate_index,
|
||||
retry_index=retry_index,
|
||||
user_id=user_id,
|
||||
api_key_id=api_key_id,
|
||||
username=username,
|
||||
api_key_name=api_key_name,
|
||||
provider_id=str(candidate.provider.id),
|
||||
endpoint_id=str(candidate.endpoint.id),
|
||||
key_id=str(candidate.key.id),
|
||||
status="skipped",
|
||||
skip_reason=skip_reason,
|
||||
is_cached=bool(getattr(candidate, "is_cached", False)),
|
||||
extra_data=extra,
|
||||
)
|
||||
# ensure visible for subsequent recorder reads
|
||||
if self.db.in_transaction():
|
||||
self.db.commit()
|
||||
return str(row.id)
|
||||
provider_id = str(candidate.provider.id)
|
||||
endpoint_id = str(candidate.endpoint.id)
|
||||
key_id = str(candidate.key.id)
|
||||
is_cached = bool(getattr(candidate, "is_cached", False))
|
||||
|
||||
def _do() -> str:
|
||||
row = RequestCandidateService.create_candidate(
|
||||
db=self.db,
|
||||
request_id=request_id,
|
||||
candidate_index=candidate_index,
|
||||
retry_index=retry_index,
|
||||
user_id=user_id,
|
||||
api_key_id=api_key_id,
|
||||
username=username,
|
||||
api_key_name=api_key_name,
|
||||
provider_id=provider_id,
|
||||
endpoint_id=endpoint_id,
|
||||
key_id=key_id,
|
||||
status="skipped",
|
||||
skip_reason=skip_reason,
|
||||
is_cached=is_cached,
|
||||
extra_data=extra,
|
||||
)
|
||||
# ensure visible for subsequent recorder reads
|
||||
if self.db.in_transaction():
|
||||
self.db.commit()
|
||||
return str(row.id)
|
||||
|
||||
return await self._db_op(_do)
|
||||
|
||||
def _build_pool_extra_data(self, candidate: ProviderCandidate) -> dict[str, Any]:
|
||||
extra: dict[str, Any] = {}
|
||||
@@ -1576,7 +1704,7 @@ class FailoverEngine:
|
||||
self._sanitize(str(inner)),
|
||||
)
|
||||
|
||||
def _mark_candidate_skipped(
|
||||
async def _mark_candidate_skipped(
|
||||
self,
|
||||
*,
|
||||
candidate_record_map: dict[tuple[int, int], str],
|
||||
@@ -1586,19 +1714,23 @@ class FailoverEngine:
|
||||
skip_reason: str | None,
|
||||
) -> None:
|
||||
max_retries = self._get_max_retries(candidate, retry_policy)
|
||||
now = datetime.now(timezone.utc)
|
||||
for retry_index in range(max_retries):
|
||||
record_id = candidate_record_map.get((candidate_index, retry_index))
|
||||
if record_id:
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="skipped",
|
||||
skip_reason=skip_reason,
|
||||
finished_at=now,
|
||||
)
|
||||
self.db.commit()
|
||||
record_ids = [
|
||||
candidate_record_map[candidate_index, ri]
|
||||
for ri in range(max_retries)
|
||||
if (candidate_index, ri) in candidate_record_map
|
||||
]
|
||||
if not record_ids:
|
||||
return
|
||||
|
||||
def _mark_remaining_slots_unused(
|
||||
def _do() -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
for rid in record_ids:
|
||||
self._update_record(rid, status="skipped", skip_reason=skip_reason, finished_at=now)
|
||||
self.db.commit()
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
async def _mark_remaining_slots_unused(
|
||||
self,
|
||||
*,
|
||||
candidate_record_map: dict[tuple[int, int], str],
|
||||
@@ -1607,7 +1739,7 @@ class FailoverEngine:
|
||||
success_retry_idx: int,
|
||||
retry_policy: RetryPolicy,
|
||||
) -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
record_ids: list[str] = []
|
||||
for candidate_idx, cand in enumerate(candidates):
|
||||
max_retries = self._get_max_retries(cand, retry_policy)
|
||||
for retry_idx in range(max_retries):
|
||||
@@ -1617,14 +1749,20 @@ class FailoverEngine:
|
||||
continue
|
||||
record_id = candidate_record_map.get((candidate_idx, retry_idx))
|
||||
if record_id:
|
||||
self._update_record(
|
||||
record_id,
|
||||
status="unused",
|
||||
finished_at=now,
|
||||
)
|
||||
self.db.commit()
|
||||
record_ids.append(record_id)
|
||||
|
||||
def _mark_candidate_remaining_retries_unused(
|
||||
if not record_ids:
|
||||
return
|
||||
|
||||
def _do() -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
for rid in record_ids:
|
||||
self._update_record(rid, status="unused", finished_at=now)
|
||||
self.db.commit()
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
async def _mark_candidate_remaining_retries_unused(
|
||||
self,
|
||||
*,
|
||||
candidate_record_map: dict[tuple[int, int], str],
|
||||
@@ -1633,21 +1771,27 @@ class FailoverEngine:
|
||||
retry_policy: RetryPolicy,
|
||||
) -> None:
|
||||
# Only meaningful for PRE_EXPAND.
|
||||
# We don't have access to candidate object list here, so infer max_retries from map keys.
|
||||
# Fallback to retry_policy.max_retries.
|
||||
now = datetime.now(timezone.utc)
|
||||
# try best-effort upper bound
|
||||
upper = max(
|
||||
(ri for (ci, ri) in candidate_record_map.keys() if ci == candidate_idx),
|
||||
default=retry_policy.max_retries - 1,
|
||||
)
|
||||
for retry_idx in range(from_retry_idx, upper + 1):
|
||||
record_id = candidate_record_map.get((candidate_idx, retry_idx))
|
||||
if record_id:
|
||||
self._update_record(record_id, status="unused", finished_at=now)
|
||||
self.db.commit()
|
||||
record_ids = [
|
||||
candidate_record_map[candidate_idx, ri]
|
||||
for ri in range(from_retry_idx, upper + 1)
|
||||
if (candidate_idx, ri) in candidate_record_map
|
||||
]
|
||||
if not record_ids:
|
||||
return
|
||||
|
||||
def _mark_retry_indices_status(
|
||||
def _do() -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
for rid in record_ids:
|
||||
self._update_record(rid, status="unused", finished_at=now)
|
||||
self.db.commit()
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
async def _mark_retry_indices_status(
|
||||
self,
|
||||
*,
|
||||
candidate_record_map: dict[tuple[int, int], str],
|
||||
@@ -1656,33 +1800,45 @@ class FailoverEngine:
|
||||
status: str,
|
||||
skip_reason: str | None = None,
|
||||
) -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
for retry_idx in retry_indices:
|
||||
record_id = candidate_record_map.get((candidate_idx, retry_idx))
|
||||
if not record_id:
|
||||
continue
|
||||
values: dict[str, Any] = {"status": status, "finished_at": now}
|
||||
if status == "skipped":
|
||||
values["skip_reason"] = skip_reason
|
||||
self._update_record(record_id, **values)
|
||||
self.db.commit()
|
||||
record_ids = [
|
||||
candidate_record_map[candidate_idx, ri]
|
||||
for ri in retry_indices
|
||||
if (candidate_idx, ri) in candidate_record_map
|
||||
]
|
||||
if not record_ids:
|
||||
return
|
||||
|
||||
def _mark_all_remaining_available_unused(
|
||||
def _do() -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
for rid in record_ids:
|
||||
values: dict[str, Any] = {"status": status, "finished_at": now}
|
||||
if status == "skipped":
|
||||
values["skip_reason"] = skip_reason
|
||||
self._update_record(rid, **values)
|
||||
self.db.commit()
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
async def _mark_all_remaining_available_unused(
|
||||
self, candidate_record_map: dict[tuple[int, int], str]
|
||||
) -> None:
|
||||
# As a safety net: do not leave available records behind in PRE_EXPAND mode.
|
||||
try:
|
||||
ids = list(candidate_record_map.values())
|
||||
if not ids:
|
||||
return
|
||||
now = datetime.now(timezone.utc)
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id.in_(ids))
|
||||
.where(RequestCandidate.status == "available")
|
||||
.values(status="unused", finished_at=now)
|
||||
)
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise
|
||||
ids = list(candidate_record_map.values())
|
||||
if not ids:
|
||||
return
|
||||
|
||||
def _do() -> None:
|
||||
try:
|
||||
now = datetime.now(timezone.utc)
|
||||
self.db.execute(
|
||||
update(RequestCandidate)
|
||||
.where(RequestCandidate.id.in_(ids))
|
||||
.where(RequestCandidate.status == "available")
|
||||
.values(status="unused", finished_at=now)
|
||||
)
|
||||
self.db.commit()
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise
|
||||
|
||||
await self._db_op(_do)
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
from decimal import Decimal
|
||||
from datetime import datetime, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -421,102 +422,108 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
|
||||
usage_params, total_cost = await cls._prepare_usage_record(params)
|
||||
total_cost = to_money_decimal(total_cost)
|
||||
|
||||
# 检查是否已存在相同 request_id 的记录
|
||||
existing_usage = (
|
||||
db.query(Usage).filter(Usage.request_id == request_id).with_for_update().first()
|
||||
)
|
||||
if existing_usage:
|
||||
if cls._is_usage_finalized(existing_usage):
|
||||
logger.debug(
|
||||
"request_id {} 已完成结算,跳过重复记账 (billing_status={})",
|
||||
request_id,
|
||||
getattr(existing_usage, "billing_status", None),
|
||||
)
|
||||
return existing_usage
|
||||
logger.debug(
|
||||
f"request_id {request_id} 已存在,更新现有记录 "
|
||||
f"(status: {existing_usage.status} -> {status})"
|
||||
def _sync_record() -> Usage:
|
||||
"""同步 DB 操作: with_for_update + 批量 update + commit"""
|
||||
from sqlalchemy import func as sql_func
|
||||
from sqlalchemy import update as sa_update
|
||||
|
||||
from src.models.database import ApiKey as ApiKeyModel
|
||||
from src.models.database import GlobalModel
|
||||
|
||||
# 检查是否已存在相同 request_id 的记录
|
||||
existing_usage = (
|
||||
db.query(Usage).filter(Usage.request_id == request_id).with_for_update().first()
|
||||
)
|
||||
cls._update_existing_usage(existing_usage, usage_params, target_model)
|
||||
usage = existing_usage
|
||||
else:
|
||||
usage = Usage(**usage_params)
|
||||
db.add(usage)
|
||||
|
||||
# 确保 user 和 api_key 在会话中
|
||||
if user and not db.object_session(user):
|
||||
user = db.merge(user)
|
||||
if api_key and not db.object_session(api_key):
|
||||
api_key = db.merge(api_key)
|
||||
|
||||
# 使用原子更新避免并发竞态条件
|
||||
from sqlalchemy import func as sql_func
|
||||
from sqlalchemy import update
|
||||
|
||||
from src.models.database import ApiKey as ApiKeyModel
|
||||
from src.models.database import GlobalModel
|
||||
|
||||
accounted, charge_applied = cls._finalize_usage_billing(
|
||||
db,
|
||||
usage=usage,
|
||||
total_cost=total_cost,
|
||||
status=status,
|
||||
)
|
||||
|
||||
if accounted:
|
||||
# 更新 API 密钥使用量
|
||||
if api_key:
|
||||
values: dict[str, Any] = {
|
||||
"total_requests": ApiKeyModel.total_requests + 1,
|
||||
"last_used_at": sql_func.now(),
|
||||
"updated_at": sql_func.now(),
|
||||
}
|
||||
if charge_applied:
|
||||
values["total_cost_usd"] = ApiKeyModel.total_cost_usd + Decimal(
|
||||
str(to_money_decimal(total_cost))
|
||||
if existing_usage:
|
||||
if cls._is_usage_finalized(existing_usage):
|
||||
logger.debug(
|
||||
"request_id {} 已完成结算,跳过重复记账 (billing_status={})",
|
||||
request_id,
|
||||
getattr(existing_usage, "billing_status", None),
|
||||
)
|
||||
db.execute(update(ApiKeyModel).where(ApiKeyModel.id == api_key.id).values(**values))
|
||||
return existing_usage
|
||||
logger.debug(
|
||||
f"request_id {request_id} 已存在,更新现有记录 "
|
||||
f"(status: {existing_usage.status} -> {status})"
|
||||
)
|
||||
cls._update_existing_usage(existing_usage, usage_params, target_model)
|
||||
usage = existing_usage
|
||||
else:
|
||||
usage = Usage(**usage_params)
|
||||
db.add(usage)
|
||||
|
||||
# 更新 GlobalModel 使用计数
|
||||
db.execute(
|
||||
update(GlobalModel)
|
||||
.where(GlobalModel.name == model)
|
||||
.values(usage_count=GlobalModel.usage_count + 1)
|
||||
# 确保 user 和 api_key 在会话中
|
||||
nonlocal user, api_key
|
||||
if user and not db.object_session(user):
|
||||
user = db.merge(user)
|
||||
if api_key and not db.object_session(api_key):
|
||||
api_key = db.merge(api_key)
|
||||
|
||||
accounted, charge_applied = cls._finalize_usage_billing(
|
||||
db,
|
||||
usage=usage,
|
||||
total_cost=total_cost,
|
||||
status=status,
|
||||
)
|
||||
|
||||
# 更新用户-模型调用次数计数器
|
||||
cls._increment_user_model_usage(db, user, model)
|
||||
if accounted:
|
||||
# 更新 API 密钥使用量
|
||||
if api_key:
|
||||
values: dict[str, Any] = {
|
||||
"total_requests": ApiKeyModel.total_requests + 1,
|
||||
"last_used_at": sql_func.now(),
|
||||
"updated_at": sql_func.now(),
|
||||
}
|
||||
if charge_applied:
|
||||
values["total_cost_usd"] = ApiKeyModel.total_cost_usd + Decimal(
|
||||
str(to_money_decimal(total_cost))
|
||||
)
|
||||
db.execute(
|
||||
sa_update(ApiKeyModel).where(ApiKeyModel.id == api_key.id).values(**values)
|
||||
)
|
||||
|
||||
# 更新 Provider 月度使用量(Provider 端真实成本,无论钱包是否扣费)
|
||||
if provider_id:
|
||||
actual_total_cost = Decimal(str(usage_params["actual_total_cost_usd"]))
|
||||
# 更新 GlobalModel 使用计数
|
||||
db.execute(
|
||||
update(Provider)
|
||||
.where(Provider.id == provider_id)
|
||||
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
||||
sa_update(GlobalModel)
|
||||
.where(GlobalModel.name == model)
|
||||
.values(usage_count=GlobalModel.usage_count + 1)
|
||||
)
|
||||
|
||||
# 更新手动代理节点请求计数(tunnel 节点由心跳上报,不在此处统计)
|
||||
manual_node_id = _extract_manual_proxy_node_id(metadata)
|
||||
if manual_node_id:
|
||||
failed = {manual_node_id: 1} if status == "failed" else None
|
||||
_increment_proxy_node_requests(db, {manual_node_id: 1}, failed)
|
||||
# 更新用户-模型调用次数计数器
|
||||
cls._increment_user_model_usage(db, user, model)
|
||||
|
||||
dispatch_codex_quota_sync_from_response_headers(
|
||||
provider_api_key_id=provider_api_key_id,
|
||||
response_headers=response_headers,
|
||||
db=db,
|
||||
)
|
||||
# 更新 Provider 月度使用量(Provider 端真实成本,无论钱包是否扣费)
|
||||
if provider_id:
|
||||
actual_total_cost = Decimal(str(usage_params["actual_total_cost_usd"]))
|
||||
db.execute(
|
||||
sa_update(Provider)
|
||||
.where(Provider.id == provider_id)
|
||||
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
|
||||
)
|
||||
|
||||
# 提交事务
|
||||
try:
|
||||
db.commit()
|
||||
except Exception as e:
|
||||
logger.error("提交使用记录时出错: {}", e)
|
||||
db.rollback()
|
||||
raise
|
||||
# 更新手动代理节点请求计数(tunnel 节点由心跳上报,不在此处统计)
|
||||
manual_node_id = _extract_manual_proxy_node_id(metadata)
|
||||
if manual_node_id:
|
||||
failed = {manual_node_id: 1} if status == "failed" else None
|
||||
_increment_proxy_node_requests(db, {manual_node_id: 1}, failed)
|
||||
|
||||
return usage
|
||||
dispatch_codex_quota_sync_from_response_headers(
|
||||
provider_api_key_id=provider_api_key_id,
|
||||
response_headers=response_headers,
|
||||
db=db,
|
||||
)
|
||||
|
||||
# 提交事务
|
||||
try:
|
||||
db.commit()
|
||||
except Exception as e:
|
||||
logger.error("提交使用记录时出错: {}", e)
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
return usage
|
||||
|
||||
return await asyncio.to_thread(_sync_record)
|
||||
|
||||
@classmethod
|
||||
async def record_usage_with_custom_cost(
|
||||
|
||||
@@ -5,6 +5,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
from collections import deque
|
||||
@@ -518,7 +519,8 @@ class StreamUsageTracker:
|
||||
# yield 后再更新数据库状态(仅第一个 chunk 时执行)
|
||||
if chunk_count == 1 and self.request_id:
|
||||
try:
|
||||
UsageService.update_usage_status(
|
||||
await asyncio.to_thread(
|
||||
UsageService.update_usage_status,
|
||||
db=self.db,
|
||||
request_id=self.request_id,
|
||||
status="streaming",
|
||||
@@ -1022,7 +1024,8 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
|
||||
# yield 后再更新数据库状态(仅第一个 chunk 时执行)
|
||||
if chunk_count == 1 and self.request_id:
|
||||
try:
|
||||
UsageService.update_usage_status(
|
||||
await asyncio.to_thread(
|
||||
UsageService.update_usage_status,
|
||||
db=self.db,
|
||||
request_id=self.request_id,
|
||||
status="streaming",
|
||||
|
||||
Reference in New Issue
Block a user