"""
AI Rate Limiter - paid AI provider usage and spend guard
=========================================
防止最終付費備援用量暴衝;超過任一硬閘即 no-write fail closed。
功能:
- 每分鐘請求限制 (RPM)
- 每日請求限制
- 每日 Token 限制
- 🔴 累積成本限制 ($5 USD) - 2026-03-29 ogt 新增
- model-aware 定價、reserved/actual/estimated receipt
- 超限阻擋 + Telegram 告警
版本: v2.0
建立日期: 2026-03-26 21:00 (台北時區)
更新日期: 2026-07-14 (台北時區)
建立者: Claude Code
"""
import hashlib
import json
import math
import re
from dataclasses import dataclass
from decimal import Decimal, InvalidOperation
from typing import Any
from uuid import uuid4
import structlog
logger = structlog.get_logger(__name__)
# =============================================================================
# Configuration - 閥值設定
# =============================================================================
RATE_LIMITS = {
"gemini": {
"rpm": 10, # 每分鐘請求數
"daily_requests": 500, # 每日請求數
"daily_tokens": 100_000, # 每日 Token 數
},
"claude": {
"rpm": 5,
"daily_requests": 200,
"daily_tokens": 50_000,
},
# 2026-03-31 ogt: NVIDIA NIM 免費版無每日限制!
# 只保留 RPM 限制 (併發控制) + 極大的 daily 上限 (監控用)
# 2026-04-03 ogt: I3 — "nvidia" → "openclaw_nemo" 對齊 AIProviderEnum (Phase 24)
"openclaw_nemo": {
"rpm": 10, # 每分鐘請求數 (放寬到 10)
"daily_requests": 99999, # 🔴 免費版無限制!設大數避免誤觸
"daily_tokens": 9999999, # 免費版無限制
},
}
# =============================================================================
# 2026-03-29 ogt: 累積成本限制 (統帥要求)
# =============================================================================
COST_LIMITS = {
"gemini": {
"total_cost_usd": 5.0, # 🔴 總成本上限 $5 USD,超過自動停用
"alert_threshold_usd": 4.0, # 警告閾值 $4 USD
},
"claude": {
"total_cost_usd": 10.0,
"alert_threshold_usd": 8.0,
},
# 2026-03-29 ogt: ADR-036 Nemotron (免費 Tier,設定低限制作為監控)
# 2026-03-31 ogt: 修復 $0.00 >= $0.00 永遠 True 的 Bug,改用大數值表示無限制
# 2026-04-03 ogt: I3 — "nvidia" → "openclaw_nemo" 對齊 AIProviderEnum (Phase 24)
"openclaw_nemo": {
"total_cost_usd": 999999.0, # 免費 Tier 無成本限制
"alert_threshold_usd": 0.0, # 不發送成本告警
},
}
@dataclass(frozen=True)
class PaidProviderPricingPolicy:
"""Reviewed exact-model pricing used by the paid-fallback guard."""
model: str
input_usd_per_million: float
output_usd_per_million: float
source: str
version: str
checked_at: str
def cost_usd(self, prompt_tokens: int, completion_tokens: int) -> float:
return (
max(0, int(prompt_tokens)) * self.input_usd_per_million
+ max(0, int(completion_tokens)) * self.output_usd_per_million
) / 1_000_000
# Unknown models are intentionally absent and therefore fail closed. Updating a
# model requires an explicit, reviewed pricing policy update in the same change.
GEMINI_PRICING_POLICIES: dict[str, PaidProviderPricingPolicy] = {
"gemini-2.5-flash-lite": PaidProviderPricingPolicy(
model="gemini-2.5-flash-lite",
input_usd_per_million=0.10,
output_usd_per_million=0.40,
source=(
"https://ai.google.dev/gemini-api/docs/pricing" "#gemini-2.5-flash-lite"
),
version="google-ai-standard-text-2026-07-14",
checked_at="2026-07-14",
),
}
# Sonnet 5 had an introductory $2/$10 price through 2026-08-31. The guard
# deliberately reserves the higher published standard $3/$15 rate so a
# promotion expiry cannot under-reserve spend without a source change.
CLAUDE_PRICING_POLICIES: dict[str, PaidProviderPricingPolicy] = {
"claude-sonnet-5": PaidProviderPricingPolicy(
model="claude-sonnet-5",
input_usd_per_million=3.0,
output_usd_per_million=15.0,
source="https://platform.claude.com/docs/en/about-claude/models/overview",
version="anthropic-standard-ceiling-2026-07-15",
checked_at="2026-07-15",
),
}
# Backward-compatible public type name used by existing Gemini tests/imports.
GeminiPricingPolicy = PaidProviderPricingPolicy
def get_gemini_pricing_policy(model: str) -> PaidProviderPricingPolicy | None:
"""Return a verified exact-model policy; no prefix or family matching."""
return GEMINI_PRICING_POLICIES.get(str(model).strip())
def get_claude_pricing_policy(model: str) -> PaidProviderPricingPolicy | None:
"""Return the reviewed exact-model Anthropic policy; aliases fail closed."""
return CLAUDE_PRICING_POLICIES.get(str(model).strip())
def get_paid_provider_pricing_policy(
provider: str,
model: str,
) -> PaidProviderPricingPolicy | None:
"""Resolve exact provider/model pricing without family or prefix matching."""
if provider == "gemini":
return get_gemini_pricing_policy(model)
if provider == "claude":
return get_claude_pricing_policy(model)
return None
# Redis Keys
REDIS_KEY_PREFIX = "ai_rate:"
RPM_KEY = f"{REDIS_KEY_PREFIX}rpm:{{provider}}"
DAILY_REQ_KEY = f"{REDIS_KEY_PREFIX}daily_req:{{provider}}:{{date}}"
DAILY_TOKEN_KEY = f"{REDIS_KEY_PREFIX}daily_token:{{provider}}:{{date}}"
DAILY_TOKEN_RESERVED_KEY = (
f"{REDIS_KEY_PREFIX}daily_token_reserved:{{provider}}:{{date}}"
)
DAILY_COST_MICRO_USD_KEY = (
f"{REDIS_KEY_PREFIX}daily_cost_micro_usd:{{provider}}:{{date}}"
)
DAILY_COST_RESERVED_MICRO_USD_KEY = (
f"{REDIS_KEY_PREFIX}daily_cost_reserved_micro_usd:{{provider}}:{{date}}"
)
TOTAL_COST_RESERVED_MICRO_USD_KEY = (
f"{REDIS_KEY_PREFIX}total_cost_reserved_micro_usd:{{provider}}"
)
GENERATION_RECEIPT_KEY = f"{REDIS_KEY_PREFIX}receipt:{{provider}}:{{receipt_id}}"
GENERATION_RUN_IDEMPOTENCY_KEY = f"{REDIS_KEY_PREFIX}run:{{provider}}:{{identity_hash}}"
GENERATION_FAILURE_COOLDOWN_KEY = (
f"{REDIS_KEY_PREFIX}failure_cooldown:{{provider}}:{{work_item_hash}}"
)
PAID_RUN_BUDGET_KEY = f"{REDIS_KEY_PREFIX}paid_run_budget:{{budget_id}}"
PROVIDER_AUTH_STATUS_KEY = f"{REDIS_KEY_PREFIX}auth_status:{{provider}}"
# 2026-03-29 ogt: 累積成本 Key (不過期,手動重置)
TOTAL_COST_KEY = f"{REDIS_KEY_PREFIX}total_cost:{{provider}}"
COST_ALERT_SENT_KEY = f"{REDIS_KEY_PREFIX}cost_alert_sent:{{provider}}"
_MICRO_USD = 1_000_000
_PAID_CANARY_WORK_ITEM_ID = "AIA-SRE-013"
_PAID_CANARY_MAX_COST_MICRO_USD = 250_000
_CORRELATION_VALUE_PATTERN = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:/-]{0,159}$")
@dataclass(frozen=True)
class AIGenerationReservation:
"""One fail-closed paid-provider reservation and its receipt identity."""
provider: str
receipt_id: str
allowed: bool
reason: str
estimated_prompt_tokens: int
estimated_completion_tokens: int
estimated_cost_usd: float
date: str
model: str = ""
pricing_source: str = ""
pricing_version: str = ""
pricing_checked_at: str = ""
input_usd_per_million: float = 0.0
output_usd_per_million: float = 0.0
trace_id: str = ""
run_id: str = ""
work_item_id: str = ""
identity_hash: str = ""
run_budget_id: str = ""
run_budget_max_micro_usd: int = 0
def actual_cost_usd(self, prompt_tokens: int, completion_tokens: int) -> float:
"""Calculate provider-reported usage with the reserved pricing policy."""
if self.input_usd_per_million <= 0 or self.output_usd_per_million <= 0:
return 0.0
return (
max(0, int(prompt_tokens)) * self.input_usd_per_million
+ max(0, int(completion_tokens)) * self.output_usd_per_million
) / 1_000_000
_RESERVE_GENERATION_LUA = r"""
local rpm_limit = tonumber(ARGV[1])
local daily_request_limit = tonumber(ARGV[2])
local daily_token_limit = tonumber(ARGV[3])
local daily_cost_limit = tonumber(ARGV[4])
local total_cost_limit = tonumber(ARGV[5])
local estimated_tokens = tonumber(ARGV[6])
local estimated_cost = tonumber(ARGV[7])
local daily_ttl = tonumber(ARGV[8])
local run_budget_required = ARGV[24] == '1'
local run_budget_id = ARGV[25]
local run_budget_max = tonumber(ARGV[26]) or 0
local function receipt(status, reason)
local reserved_prompt_tokens = '0'
local reserved_completion_tokens = '0'
local reserved_total_tokens = '0'
local reserved_cost_micro_usd = '0'
local run_budget_reserved_cost_micro_usd = '0'
local terminal_state = 'blocked_no_write'
if status == 'reserved' then
reserved_prompt_tokens = ARGV[16]
reserved_completion_tokens = ARGV[17]
reserved_total_tokens = ARGV[6]
reserved_cost_micro_usd = ARGV[7]
if run_budget_required then
run_budget_reserved_cost_micro_usd = ARGV[7]
end
terminal_state = 'pending_finalize'
end
redis.call('HSET', KEYS[9],
'schema_version', 'ai_paid_generation_receipt_v1',
'provider', ARGV[9],
'receipt_id', ARGV[10],
'status', status,
'reason', reason,
'terminal_state', terminal_state,
'created_at', ARGV[11],
'date', ARGV[12],
'trace_id', ARGV[13],
'run_id', ARGV[14],
'work_item_id', ARGV[15],
'model', ARGV[18],
'pricing_source', ARGV[19],
'pricing_version', ARGV[20],
'pricing_checked_at', ARGV[21],
'pricing_input_usd_per_million', ARGV[22],
'pricing_output_usd_per_million', ARGV[23],
'run_budget_required', ARGV[24],
'run_budget_id', run_budget_id,
'run_budget_max_micro_usd', ARGV[26],
'run_budget_reserved_cost_micro_usd', run_budget_reserved_cost_micro_usd,
'estimated_prompt_tokens', ARGV[16],
'estimated_completion_tokens', ARGV[17],
'estimated_total_tokens', ARGV[6],
'estimated_cost_micro_usd', ARGV[7],
'reserved_prompt_tokens', reserved_prompt_tokens,
'reserved_completion_tokens', reserved_completion_tokens,
'reserved_total_tokens', reserved_total_tokens,
'reserved_cost_micro_usd', reserved_cost_micro_usd,
'actual_prompt_tokens', '0',
'actual_completion_tokens', '0',
'actual_total_tokens', '0',
'actual_cost_micro_usd', '0')
end
-- One paid provider call is permitted for one canonical run_id; trace_id and
-- work_item_id remain mandatory receipt correlation. The claim and the
-- receipt/budget transition live in the same Lua execution so concurrent
-- workers cannot double-reserve or double-call a paid provider.
local existing_receipt_id = redis.call('GET', KEYS[10])
if existing_receipt_id then
return {0, 'idempotent_replay', existing_receipt_id}
end
local claimed = redis.call('SET', KEYS[10], ARGV[10], 'NX')
if not claimed then
return {0, 'idempotent_replay', redis.call('GET', KEYS[10]) or ARGV[10]}
end
if redis.call('EXISTS', KEYS[11]) == 1 then
receipt('blocked', 'failure_cooldown')
return {0, 'failure_cooldown', ARGV[10]}
end
local current_rpm = tonumber(redis.call('GET', KEYS[1]) or '0')
local current_requests = tonumber(redis.call('GET', KEYS[2]) or '0')
local current_tokens = tonumber(redis.call('GET', KEYS[3]) or '0')
local reserved_tokens = tonumber(redis.call('GET', KEYS[4]) or '0')
local current_daily_cost = tonumber(redis.call('GET', KEYS[5]) or '0')
local reserved_daily_cost = tonumber(redis.call('GET', KEYS[6]) or '0')
local current_total_cost = math.floor(
(tonumber(redis.call('GET', KEYS[7]) or '0') * 1000000) + 0.5
)
local reserved_total_cost = tonumber(redis.call('GET', KEYS[8]) or '0')
local run_budget_reserved = 0
local run_budget_accounted = 0
if run_budget_required then
if run_budget_id == '' or run_budget_max <= 0 then
receipt('blocked', 'run_cost_cap_invalid')
return {0, 'run_cost_cap_invalid', ARGV[10]}
end
local existing_budget_id = redis.call('HGET', KEYS[12], 'run_budget_id')
local existing_budget_max = tonumber(
redis.call('HGET', KEYS[12], 'max_cost_micro_usd') or '0'
)
if existing_budget_id and existing_budget_id ~= run_budget_id then
receipt('blocked', 'run_cost_cap_mismatch')
return {0, 'run_cost_cap_mismatch', ARGV[10]}
end
if existing_budget_max > 0 and existing_budget_max ~= run_budget_max then
receipt('blocked', 'run_cost_cap_mismatch')
return {0, 'run_cost_cap_mismatch', ARGV[10]}
end
run_budget_reserved = tonumber(
redis.call('HGET', KEYS[12], 'reserved_cost_micro_usd') or '0'
)
run_budget_accounted = tonumber(
redis.call('HGET', KEYS[12], 'accounted_cost_micro_usd') or '0'
)
if run_budget_accounted + run_budget_reserved + estimated_cost > run_budget_max then
receipt('blocked', 'run_cost_limit')
return {0, 'run_cost_limit', ARGV[10]}
end
end
if current_rpm + 1 > rpm_limit then
receipt('blocked', 'rpm_limit')
return {0, 'rpm_limit', ARGV[10]}
end
if current_requests + 1 > daily_request_limit then
receipt('blocked', 'daily_request_limit')
return {0, 'daily_request_limit', ARGV[10]}
end
if current_tokens + reserved_tokens + estimated_tokens > daily_token_limit then
receipt('blocked', 'daily_token_limit')
return {0, 'daily_token_limit', ARGV[10]}
end
if current_daily_cost + reserved_daily_cost + estimated_cost > daily_cost_limit then
receipt('blocked', 'daily_cost_limit')
return {0, 'daily_cost_limit', ARGV[10]}
end
if current_total_cost + reserved_total_cost + estimated_cost > total_cost_limit then
receipt('blocked', 'total_cost_limit')
return {0, 'total_cost_limit', ARGV[10]}
end
local new_rpm = redis.call('INCR', KEYS[1])
if new_rpm == 1 or redis.call('TTL', KEYS[1]) < 0 then redis.call('EXPIRE', KEYS[1], 60) end
local new_requests = redis.call('INCR', KEYS[2])
if new_requests == 1 or redis.call('TTL', KEYS[2]) < 0 then
redis.call('EXPIRE', KEYS[2], daily_ttl)
end
redis.call('INCRBY', KEYS[4], estimated_tokens)
redis.call('EXPIRE', KEYS[4], daily_ttl)
redis.call('INCRBY', KEYS[6], estimated_cost)
redis.call('EXPIRE', KEYS[6], daily_ttl)
-- Cumulative paid-spend reservations do not expire. A missing finalize receipt
-- must keep consuming the hard cap until reconciliation/reset, never turn green.
redis.call('INCRBY', KEYS[8], estimated_cost)
if run_budget_required then
redis.call('HSET', KEYS[12],
'schema_version', 'ai_paid_run_budget_receipt_v1',
'run_budget_id', run_budget_id,
'run_id', ARGV[14],
'work_item_id', ARGV[15],
'max_cost_micro_usd', run_budget_max,
'status', 'pending_finalize')
redis.call('HSETNX', KEYS[12], 'reserved_cost_micro_usd', '0')
redis.call('HSETNX', KEYS[12], 'accounted_cost_micro_usd', '0')
redis.call('HSETNX', KEYS[12], 'reservation_count', '0')
redis.call('HSETNX', KEYS[12], 'finalized_count', '0')
redis.call('HINCRBY', KEYS[12], 'reserved_cost_micro_usd', estimated_cost)
redis.call('HINCRBY', KEYS[12], 'reservation_count', 1)
end
receipt('reserved', 'allowed')
return {1, 'allowed', ARGV[10]}
"""
_WRITE_BLOCKED_GENERATION_RECEIPT_LUA = r"""
local receipt_id = ARGV[1]
local payload = cjson.decode(ARGV[2])
local has_identity = ARGV[3] == '1'
local existing_receipt_id = has_identity and redis.call('GET', KEYS[2]) or nil
if existing_receipt_id then
return {0, 'idempotent_replay', existing_receipt_id}
end
if redis.call('EXISTS', KEYS[1]) == 1 then
if has_identity then redis.call('SET', KEYS[2], receipt_id, 'NX') end
return {0, 'receipt_exists', redis.call('HGET', KEYS[1], 'receipt_id') or receipt_id}
end
if has_identity then
local claimed = redis.call('SET', KEYS[2], receipt_id, 'NX')
if not claimed then
return {0, 'idempotent_replay', redis.call('GET', KEYS[2]) or receipt_id}
end
end
for field, value in pairs(payload) do
redis.call('HSET', KEYS[1], field, tostring(value))
end
return {1, 'written', receipt_id}
"""
_FINALIZE_GENERATION_LUA = r"""
local status = redis.call('HGET', KEYS[7], 'status')
if not status then return {0, 'receipt_missing'} end
if status ~= 'reserved' then return {2, status} end
local estimated_tokens = tonumber(redis.call('HGET', KEYS[7], 'estimated_total_tokens') or '0')
local estimated_cost = tonumber(redis.call('HGET', KEYS[7], 'estimated_cost_micro_usd') or '0')
local actual_prompt_tokens = tonumber(ARGV[1])
local actual_completion_tokens = tonumber(ARGV[2])
local actual_tokens = actual_prompt_tokens + actual_completion_tokens
local actual_cost = tonumber(ARGV[3])
local success = ARGV[4] == '1'
local accounted_tokens = actual_tokens
local accounted_cost = actual_cost
local usage_source = 'provider_metadata'
local final_status = 'succeeded'
if not success then
accounted_tokens = estimated_tokens
accounted_cost = estimated_cost
usage_source = 'reservation_estimate'
final_status = 'failed_estimate_charged'
elseif actual_tokens <= 0 or actual_cost <= 0 then
accounted_tokens = estimated_tokens
accounted_cost = estimated_cost
usage_source = 'reservation_estimate'
final_status = 'succeeded_usage_estimated'
end
local run_budget_required = redis.call('HGET', KEYS[7], 'run_budget_required') == '1'
local run_budget_id = redis.call('HGET', KEYS[7], 'run_budget_id') or ''
local run_budget_max = tonumber(
redis.call('HGET', KEYS[7], 'run_budget_max_micro_usd') or '0'
)
local run_budget_reserved = 0
local run_budget_accounted = 0
if run_budget_required then
if run_budget_id == '' or run_budget_max <= 0 then
return {0, 'run_budget_receipt_invalid'}
end
if redis.call('HGET', KEYS[10], 'run_budget_id') ~= run_budget_id then
return {0, 'run_budget_identity_mismatch'}
end
if tonumber(redis.call('HGET', KEYS[10], 'max_cost_micro_usd') or '0') ~= run_budget_max then
return {0, 'run_budget_cap_mismatch'}
end
run_budget_reserved = tonumber(
redis.call('HGET', KEYS[10], 'reserved_cost_micro_usd') or '-1'
)
run_budget_accounted = tonumber(
redis.call('HGET', KEYS[10], 'accounted_cost_micro_usd') or '-1'
)
if run_budget_reserved < estimated_cost or run_budget_accounted < 0 then
return {0, 'run_budget_reservation_mismatch'}
end
end
local token_reserved = tonumber(redis.call('GET', KEYS[2]) or '0')
redis.call('SET', KEYS[2], math.max(0, token_reserved - estimated_tokens), 'EX', ARGV[8])
local daily_cost_reserved = tonumber(redis.call('GET', KEYS[4]) or '0')
redis.call('SET', KEYS[4], math.max(0, daily_cost_reserved - estimated_cost), 'EX', ARGV[8])
local total_cost_reserved = tonumber(redis.call('GET', KEYS[6]) or '0')
local remaining_total_reserved = math.max(0, total_cost_reserved - estimated_cost)
if remaining_total_reserved > 0 then
redis.call('SET', KEYS[6], remaining_total_reserved)
else
redis.call('DEL', KEYS[6])
end
redis.call('INCRBY', KEYS[1], accounted_tokens)
redis.call('EXPIRE', KEYS[1], ARGV[8])
redis.call('INCRBY', KEYS[3], accounted_cost)
redis.call('EXPIRE', KEYS[3], ARGV[8])
redis.call('INCRBYFLOAT', KEYS[5], accounted_cost / 1000000)
redis.call('HSET', KEYS[7],
'status', final_status,
'terminal_state', 'finalized_accounted',
'finished_at', ARGV[5],
'success', ARGV[4],
'error_code', ARGV[6],
'actual_prompt_tokens', ARGV[1],
'actual_completion_tokens', ARGV[2],
'actual_total_tokens', actual_tokens,
'actual_cost_micro_usd', ARGV[3],
'reserved_prompt_tokens', '0',
'reserved_completion_tokens', '0',
'reserved_total_tokens', '0',
'reserved_cost_micro_usd', '0',
'run_budget_reserved_cost_micro_usd', '0',
'run_budget_accounted_cost_micro_usd', accounted_cost,
'accounted_tokens', accounted_tokens,
'accounted_cost_micro_usd', accounted_cost,
'usage_source', usage_source)
if run_budget_required then
local remaining_run_reserved = run_budget_reserved - estimated_cost
local new_run_accounted = run_budget_accounted + accounted_cost
local run_budget_status = 'pending_finalize'
if remaining_run_reserved == 0 then run_budget_status = 'finalized' end
if new_run_accounted > run_budget_max then run_budget_status = 'cap_exceeded' end
redis.call('HSET', KEYS[10],
'reserved_cost_micro_usd', remaining_run_reserved,
'accounted_cost_micro_usd', new_run_accounted,
'status', run_budget_status,
'updated_at', ARGV[5])
redis.call('HINCRBY', KEYS[10], 'finalized_count', 1)
end
local authentication_state = 'unknown'
local verification_status = 'not_verified_last_generation_failed'
if success then
authentication_state = 'true'
verification_status = 'authenticated_verified_by_generation'
elseif ARGV[6] == 'authentication_rejected' then
authentication_state = 'false'
verification_status = 'authentication_rejected'
end
redis.call('HSET', KEYS[8],
'schema_version', 'ai_provider_authentication_status_v1',
'provider', ARGV[10],
'authenticated', authentication_state,
'authentication_verified', success and 'true' or 'false',
'verification_status', verification_status,
'verification_source', 'durable_generation_receipt',
'verified_at', ARGV[5],
'receipt_id', redis.call('HGET', KEYS[7], 'receipt_id') or '')
redis.call('EXPIRE', KEYS[8], ARGV[7])
if not success then
redis.call('SET', KEYS[9], redis.call('HGET', KEYS[7], 'receipt_id') or '',
'EX', ARGV[9])
end
return {1, final_status, accounted_tokens, accounted_cost}
"""
# =============================================================================
# Rate Limiter
# =============================================================================
class AIRateLimiter:
"""
AI API 用量限制器
使用 Redis 計數器追蹤用量,超限時返回降級建議。
Usage:
limiter = AIRateLimiter()
allowed, reason = await limiter.check_and_increment("gemini")
if not allowed:
# 降級到 Ollama
provider = "ollama"
"""
def __init__(self) -> None:
self._redis = None
async def _get_redis(self):
"""Lazy load Redis"""
if self._redis is None:
from src.core.redis_client import get_redis
self._redis = get_redis()
return self._redis
def _get_today(self) -> str:
"""取得今日日期 (台北時區)"""
from src.utils.timezone import now_taipei
return now_taipei().strftime("%Y-%m-%d")
def _seconds_until_tomorrow(self) -> int:
"""Expire daily counters at the next Asia/Taipei day boundary."""
from datetime import timedelta
from src.utils.timezone import now_taipei
now = now_taipei()
tomorrow = (now + timedelta(days=1)).replace(
hour=0,
minute=0,
second=0,
microsecond=0,
)
return max(60, int((tomorrow - now).total_seconds()))
@staticmethod
def _safe_receipt_context(context: dict[str, Any] | None) -> tuple[str, str, str]:
"""Keep three allowlisted, bounded ASCII correlation identifiers only."""
payload = context or {}
def _value(key: str) -> str:
value = payload.get(key)
if value is None:
return ""
candidate = str(value)
if not _CORRELATION_VALUE_PATTERN.fullmatch(candidate):
logger.warning(
"ai_generation_receipt_correlation_rejected",
field=key,
value_length=len(candidate),
)
return ""
return candidate
return (
_value("trace_id"),
_value("run_id"),
_value("work_item_id"),
)
@staticmethod
def _identity_hash(*parts: str) -> str:
"""Return a bounded, non-reversible Redis-key identity."""
payload = "\x1f".join(parts).encode("utf-8")
return hashlib.sha256(payload).hexdigest()
@classmethod
def _paid_run_budget_contract(
cls,
context: dict[str, Any] | None,
*,
run_id: str,
work_item_id: str,
) -> tuple[bool, str, int, str | None]:
"""Return one bounded aggregate paid-run budget or a fail-closed reason."""
payload = context or {}
synthetic_validation = payload.get("synthetic_validation")
canary_work_item = work_item_id == _PAID_CANARY_WORK_ITEM_ID
required = canary_work_item or "paid_canary_max_cost_usd" in payload
if not required:
return False, "", 0, None
budget_hash = cls._identity_hash(run_id, work_item_id)
budget_id = f"paid-run-{budget_hash[:32]}"
if canary_work_item and synthetic_validation is not True:
return True, budget_id, 0, "run_cost_cap_context_invalid"
raw_cap = payload.get("paid_canary_max_cost_usd")
if raw_cap is None:
return True, budget_id, 0, "run_cost_cap_missing"
if isinstance(raw_cap, bool):
return True, budget_id, 0, "run_cost_cap_invalid"
try:
cap_usd = Decimal(str(raw_cap))
except (InvalidOperation, ValueError):
return True, budget_id, 0, "run_cost_cap_invalid"
if not cap_usd.is_finite() or cap_usd <= 0 or cap_usd > Decimal("0.25"):
return True, budget_id, 0, "run_cost_cap_invalid"
cap_micro = cap_usd * _MICRO_USD
if cap_micro != cap_micro.to_integral_value():
return True, budget_id, 0, "run_cost_cap_invalid"
max_cost_micro_usd = int(cap_micro)
if not 1 <= max_cost_micro_usd <= _PAID_CANARY_MAX_COST_MICRO_USD:
return True, budget_id, 0, "run_cost_cap_invalid"
return True, budget_id, max_cost_micro_usd, None
@staticmethod
def _estimate_paid_usage(
prompt: str,
max_output_tokens: int,
pricing: PaidProviderPricingPolicy,
) -> tuple[int, int, float]:
"""Conservatively reserve paid-provider tokens and spend before generation."""
# One token per UTF-8 byte deliberately over-reserves mixed zh-TW/English
# prompts. Completion reserves the configured hard output ceiling.
prompt_tokens = max(1, len(prompt.encode("utf-8")))
completion_tokens = max(1, int(max_output_tokens))
estimated_cost = pricing.cost_usd(prompt_tokens, completion_tokens)
return prompt_tokens, completion_tokens, estimated_cost
@staticmethod
def _provider_limits(provider: str) -> tuple[dict[str, int], dict[str, float]]:
"""Use one settings-backed paid-provider quota source across runtime/readback."""
limits = dict(RATE_LIMITS[provider])
costs = dict(COST_LIMITS.get(provider, {}))
if provider not in {"claude", "gemini"}:
return limits, costs
from src.core.config import settings
prefix = provider.upper()
limits.update(
rpm=int(getattr(settings, f"{prefix}_RPM_LIMIT")),
daily_requests=int(getattr(settings, f"{prefix}_DAILY_QUOTA")),
daily_tokens=int(getattr(settings, f"{prefix}_DAILY_TOKEN_LIMIT")),
)
costs.update(
daily_cost_usd=float(getattr(settings, f"{prefix}_DAILY_COST_LIMIT_USD")),
total_cost_usd=float(getattr(settings, f"{prefix}_TOTAL_COST_LIMIT_USD")),
alert_threshold_usd=float(
getattr(settings, f"{prefix}_COST_ALERT_THRESHOLD_USD")
),
)
return limits, costs
async def _write_blocked_generation_receipt(
self,
*,
provider: str,
receipt_id: str,
reason: str,
date: str,
model: str,
context: dict[str, Any] | None,
run_budget_required: bool = False,
run_budget_id: str = "",
run_budget_max_micro_usd: int = 0,
) -> bool:
"""Best-effort durable no-reservation receipt for pre-pricing blocks."""
from src.utils.timezone import now_taipei_iso
safe_model = (
model
if re.fullmatch(r"[A-Za-z0-9][A-Za-z0-9._:-]{0,119}", model)
else "invalid_model_identifier"
)
trace_id, run_id, work_item_id = self._safe_receipt_context(context)
receipt_key = GENERATION_RECEIPT_KEY.format(
provider=provider,
receipt_id=receipt_id,
)
mapping = {
"schema_version": "ai_paid_generation_receipt_v1",
"provider": provider,
"receipt_id": receipt_id,
"status": "blocked",
"reason": reason,
"terminal_state": "blocked_no_write",
"created_at": now_taipei_iso(),
"date": date,
"trace_id": trace_id,
"run_id": run_id,
"work_item_id": work_item_id,
"model": safe_model,
"run_budget_required": "1" if run_budget_required else "0",
"run_budget_id": run_budget_id,
"run_budget_max_micro_usd": str(run_budget_max_micro_usd),
"run_budget_reserved_cost_micro_usd": "0",
"estimated_prompt_tokens": "0",
"estimated_completion_tokens": "0",
"estimated_total_tokens": "0",
"estimated_cost_micro_usd": "0",
"reserved_prompt_tokens": "0",
"reserved_completion_tokens": "0",
"reserved_total_tokens": "0",
"reserved_cost_micro_usd": "0",
"actual_prompt_tokens": "0",
"actual_completion_tokens": "0",
"actual_total_tokens": "0",
"actual_cost_micro_usd": "0",
}
try:
redis = await self._get_redis()
if redis is None:
return False
has_identity = bool(trace_id and run_id and work_item_id)
identity_hash = self._identity_hash(provider, run_id)
idempotency_key = GENERATION_RUN_IDEMPOTENCY_KEY.format(
provider=provider,
identity_hash=identity_hash,
)
result = await redis.eval(
_WRITE_BLOCKED_GENERATION_RECEIPT_LUA,
2,
receipt_key,
idempotency_key,
receipt_id,
json.dumps(mapping, sort_keys=True, separators=(",", ":")),
1 if has_identity else 0,
)
return int(result[0]) in {0, 1}
except Exception as exc:
logger.error(
"ai_generation_blocked_receipt_write_failed",
provider=provider,
receipt_id=receipt_id,
reason=reason,
error_type=type(exc).__name__,
)
return False
async def reserve_generation(
self,
provider: str,
prompt: str,
*,
model: str | None = None,
max_output_tokens: int | None = None,
context: dict[str, Any] | None = None,
) -> AIGenerationReservation:
"""Atomically reserve RPM, requests, tokens, and cost before a paid call.
Redis or configuration failure denies the call. A blocked attempt still
gets a non-secret receipt when Redis is available.
"""
if provider not in {"claude", "gemini"}:
raise ValueError("reserve_generation requires a supported paid provider")
from src.core.config import settings
from src.utils.timezone import now_taipei_iso
today = self._get_today()
provider_prefix = provider.upper()
configured_model = getattr(settings, f"{provider_prefix}_MODEL")
model_name = str(model if model is not None else configured_model).strip()
trace_id, run_id, work_item_id = self._safe_receipt_context(context)
if not (trace_id and run_id and work_item_id):
receipt_id = str(uuid4())
receipt_written = await self._write_blocked_generation_receipt(
provider=provider,
receipt_id=receipt_id,
reason="missing_correlation_identity",
date=today,
model=model_name,
context=context,
)
logger.error(
"ai_generation_reservation_failed_closed",
provider=provider,
receipt_id=receipt_id,
reason="missing_correlation_identity",
receipt_written=receipt_written,
)
return AIGenerationReservation(
provider=provider,
receipt_id=receipt_id,
allowed=False,
reason="missing_correlation_identity",
estimated_prompt_tokens=0,
estimated_completion_tokens=0,
estimated_cost_usd=0.0,
date=today,
model=model_name,
trace_id=trace_id,
run_id=run_id,
work_item_id=work_item_id,
)
identity_hash = self._identity_hash(
provider,
run_id,
)
receipt_id = f"gen-{identity_hash[:32]}"
(
run_budget_required,
run_budget_id,
run_budget_max_micro_usd,
run_budget_error,
) = self._paid_run_budget_contract(
context,
run_id=run_id,
work_item_id=work_item_id,
)
if run_budget_error:
receipt_written = await self._write_blocked_generation_receipt(
provider=provider,
receipt_id=receipt_id,
reason=run_budget_error,
date=today,
model=model_name,
context=context,
run_budget_required=run_budget_required,
run_budget_id=run_budget_id,
run_budget_max_micro_usd=run_budget_max_micro_usd,
)
logger.error(
"ai_generation_reservation_failed_closed",
provider=provider,
receipt_id=receipt_id,
reason=run_budget_error,
receipt_written=receipt_written,
)
return AIGenerationReservation(
provider=provider,
receipt_id=receipt_id,
allowed=False,
reason=run_budget_error,
estimated_prompt_tokens=0,
estimated_completion_tokens=0,
estimated_cost_usd=0.0,
date=today,
model=model_name,
trace_id=trace_id,
run_id=run_id,
work_item_id=work_item_id,
identity_hash=identity_hash,
run_budget_id=run_budget_id,
run_budget_max_micro_usd=run_budget_max_micro_usd,
)
pricing = get_paid_provider_pricing_policy(provider, model_name)
if pricing is None:
receipt_written = await self._write_blocked_generation_receipt(
provider=provider,
receipt_id=receipt_id,
reason="pricing_policy_missing",
date=today,
model=model_name,
context=context,
run_budget_required=run_budget_required,
run_budget_id=run_budget_id,
run_budget_max_micro_usd=run_budget_max_micro_usd,
)
logger.error(
"ai_generation_reservation_failed_closed",
provider=provider,
receipt_id=receipt_id,
model=model_name,
reason="pricing_policy_missing",
receipt_written=receipt_written,
)
return AIGenerationReservation(
provider=provider,
receipt_id=receipt_id,
allowed=False,
reason="pricing_policy_missing",
estimated_prompt_tokens=0,
estimated_completion_tokens=0,
estimated_cost_usd=0.0,
date=today,
model=model_name,
trace_id=trace_id,
run_id=run_id,
work_item_id=work_item_id,
identity_hash=identity_hash,
run_budget_id=run_budget_id,
run_budget_max_micro_usd=run_budget_max_micro_usd,
)
output_limit = int(
max_output_tokens
if max_output_tokens is not None
else getattr(settings, f"{provider_prefix}_MAX_OUTPUT_TOKENS")
)
prompt_tokens, completion_tokens, estimated_cost = self._estimate_paid_usage(
prompt,
output_limit,
pricing,
)
reservation = AIGenerationReservation(
provider=provider,
receipt_id=receipt_id,
allowed=False,
reason="cost_guard_unavailable",
estimated_prompt_tokens=prompt_tokens,
estimated_completion_tokens=completion_tokens,
estimated_cost_usd=estimated_cost,
date=today,
model=model_name,
pricing_source=pricing.source,
pricing_version=pricing.version,
pricing_checked_at=pricing.checked_at,
input_usd_per_million=pricing.input_usd_per_million,
output_usd_per_million=pricing.output_usd_per_million,
trace_id=trace_id,
run_id=run_id,
work_item_id=work_item_id,
identity_hash=identity_hash,
run_budget_id=run_budget_id,
run_budget_max_micro_usd=run_budget_max_micro_usd,
)
try:
limits, costs = self._provider_limits(provider)
if (
limits["rpm"] <= 0
or limits["daily_requests"] <= 0
or limits["daily_tokens"] <= 0
or costs["daily_cost_usd"] <= 0
or costs["total_cost_usd"] <= 0
or costs["alert_threshold_usd"] <= 0
or costs["alert_threshold_usd"] > costs["total_cost_usd"]
):
raise ValueError(f"invalid_{provider}_cost_guard_limits")
redis = await self._get_redis()
if redis is None:
raise RuntimeError("redis_unavailable")
daily_ttl = self._seconds_until_tomorrow()
estimated_tokens = prompt_tokens + completion_tokens
estimated_cost_micro = max(1, math.ceil(estimated_cost * _MICRO_USD))
receipt_key = GENERATION_RECEIPT_KEY.format(
provider=provider,
receipt_id=receipt_id,
)
keys = [
RPM_KEY.format(provider=provider),
DAILY_REQ_KEY.format(provider=provider, date=today),
DAILY_TOKEN_KEY.format(provider=provider, date=today),
DAILY_TOKEN_RESERVED_KEY.format(provider=provider, date=today),
DAILY_COST_MICRO_USD_KEY.format(provider=provider, date=today),
DAILY_COST_RESERVED_MICRO_USD_KEY.format(provider=provider, date=today),
TOTAL_COST_KEY.format(provider=provider),
TOTAL_COST_RESERVED_MICRO_USD_KEY.format(provider=provider),
receipt_key,
GENERATION_RUN_IDEMPOTENCY_KEY.format(
provider=provider,
identity_hash=identity_hash,
),
GENERATION_FAILURE_COOLDOWN_KEY.format(
provider=provider,
work_item_hash=self._identity_hash(provider, work_item_id),
),
PAID_RUN_BUDGET_KEY.format(
budget_id=(run_budget_id or f"unused-{identity_hash[:32]}")
),
]
result = await redis.eval(
_RESERVE_GENERATION_LUA,
len(keys),
*keys,
limits["rpm"],
limits["daily_requests"],
limits["daily_tokens"],
math.floor(costs["daily_cost_usd"] * _MICRO_USD),
math.floor(costs["total_cost_usd"] * _MICRO_USD),
estimated_tokens,
estimated_cost_micro,
daily_ttl,
provider,
receipt_id,
now_taipei_iso(),
today,
trace_id,
run_id,
work_item_id,
prompt_tokens,
completion_tokens,
model_name,
pricing.source,
pricing.version,
pricing.checked_at,
pricing.input_usd_per_million,
pricing.output_usd_per_million,
1 if run_budget_required else 0,
run_budget_id,
run_budget_max_micro_usd,
)
allowed = bool(int(result[0]))
raw_reason = result[1]
reason = (
raw_reason.decode()
if isinstance(raw_reason, bytes)
else str(raw_reason)
)
raw_receipt_id = result[2] if len(result) > 2 else receipt_id
result_receipt_id = (
raw_receipt_id.decode()
if isinstance(raw_receipt_id, bytes)
else str(raw_receipt_id)
)
reservation = AIGenerationReservation(
provider=provider,
receipt_id=result_receipt_id,
allowed=allowed,
reason=reason,
estimated_prompt_tokens=prompt_tokens,
estimated_completion_tokens=completion_tokens,
estimated_cost_usd=estimated_cost,
date=today,
model=model_name,
pricing_source=pricing.source,
pricing_version=pricing.version,
pricing_checked_at=pricing.checked_at,
input_usd_per_million=pricing.input_usd_per_million,
output_usd_per_million=pricing.output_usd_per_million,
trace_id=trace_id,
run_id=run_id,
work_item_id=work_item_id,
identity_hash=identity_hash,
run_budget_id=run_budget_id,
run_budget_max_micro_usd=run_budget_max_micro_usd,
)
logger.info(
"ai_generation_reservation",
provider=provider,
receipt_id=result_receipt_id,
allowed=allowed,
reason=reason,
estimated_tokens=estimated_tokens,
estimated_cost_usd=round(estimated_cost, 6),
model=model_name,
pricing_version=pricing.version,
)
await self._warn_for_cost_exposure(provider)
return reservation
except Exception as exc:
logger.error(
"ai_generation_reservation_failed_closed",
provider=provider,
receipt_id=receipt_id,
error_type=type(exc).__name__,
)
return reservation
async def _warn_for_cost_exposure(self, provider: str) -> None:
"""Include pending reservations when evaluating the cost alert threshold."""
try:
stats = await self.get_usage_stats(provider)
total = stats.get("total_cost_usd") or {}
exposure = float(total.get("current", 0.0)) + float(
total.get("reserved", 0.0)
)
threshold = float(total.get("alert_threshold", 0.0))
if threshold > 0 and exposure >= threshold:
await self._send_cost_warning(provider, exposure, threshold)
except Exception as exc:
logger.warning(
"ai_cost_exposure_warning_readback_failed",
provider=provider,
error_type=type(exc).__name__,
)
async def finalize_generation(
self,
reservation: AIGenerationReservation,
*,
success: bool,
prompt_tokens: int = 0,
completion_tokens: int = 0,
error_code: str = "",
) -> bool:
"""Finalize one reservation exactly once and retain a durable receipt.
Failed calls and responses without usage metadata are charged at the
conservative reservation estimate, preventing zero-token/cost false green.
"""
if not reservation.allowed:
return False
from src.core.config import settings
from src.utils.timezone import now_taipei_iso
try:
pricing = get_paid_provider_pricing_policy(
reservation.provider,
reservation.model,
)
if (
pricing is None
or reservation.pricing_source != pricing.source
or reservation.pricing_version != pricing.version
or reservation.pricing_checked_at != pricing.checked_at
or reservation.input_usd_per_million != pricing.input_usd_per_million
or reservation.output_usd_per_million != pricing.output_usd_per_million
):
raise ValueError("paid_provider_pricing_reservation_mismatch")
redis = await self._get_redis()
if redis is None:
raise RuntimeError("redis_unavailable")
actual_prompt_tokens = max(0, int(prompt_tokens))
actual_completion_tokens = max(0, int(completion_tokens))
actual_tokens = actual_prompt_tokens + actual_completion_tokens
actual_cost_usd = reservation.actual_cost_usd(
actual_prompt_tokens,
actual_completion_tokens,
)
actual_cost_micro = (
max(1, math.ceil(actual_cost_usd * _MICRO_USD))
if actual_tokens > 0
else 0
)
daily_ttl = self._seconds_until_tomorrow()
provider_prefix = reservation.provider.upper()
auth_status_ttl = int(
getattr(settings, f"{provider_prefix}_USAGE_RECEIPT_TTL_SECONDS")
)
receipt_key = GENERATION_RECEIPT_KEY.format(
provider=reservation.provider,
receipt_id=reservation.receipt_id,
)
keys = [
DAILY_TOKEN_KEY.format(
provider=reservation.provider, date=reservation.date
),
DAILY_TOKEN_RESERVED_KEY.format(
provider=reservation.provider, date=reservation.date
),
DAILY_COST_MICRO_USD_KEY.format(
provider=reservation.provider, date=reservation.date
),
DAILY_COST_RESERVED_MICRO_USD_KEY.format(
provider=reservation.provider,
date=reservation.date,
),
TOTAL_COST_KEY.format(provider=reservation.provider),
TOTAL_COST_RESERVED_MICRO_USD_KEY.format(provider=reservation.provider),
receipt_key,
PROVIDER_AUTH_STATUS_KEY.format(provider=reservation.provider),
GENERATION_FAILURE_COOLDOWN_KEY.format(
provider=reservation.provider,
work_item_hash=self._identity_hash(
reservation.provider,
reservation.work_item_id or reservation.receipt_id,
),
),
PAID_RUN_BUDGET_KEY.format(
budget_id=(
reservation.run_budget_id
or f"unused-{reservation.identity_hash[:32]}"
)
),
]
result = await redis.eval(
_FINALIZE_GENERATION_LUA,
len(keys),
*keys,
actual_prompt_tokens,
actual_completion_tokens,
actual_cost_micro,
1 if success else 0,
now_taipei_iso(),
str(error_code)[:120],
auth_status_ttl,
daily_ttl,
int(
getattr(
settings,
f"{provider_prefix}_FAILURE_COOLDOWN_SECONDS",
)
),
reservation.provider,
)
result_code = int(result[0])
raw_status = result[1]
status = (
raw_status.decode()
if isinstance(raw_status, bytes)
else str(raw_status)
)
finalized = result_code in (1, 2)
logger.info(
"ai_generation_receipt_finalized",
provider=reservation.provider,
receipt_id=reservation.receipt_id,
status=status,
finalized=finalized,
)
if result_code == 1:
total_raw = await redis.get(
TOTAL_COST_KEY.format(provider=reservation.provider)
)
total_cost = float(total_raw or 0.0)
_, costs = self._provider_limits(reservation.provider)
if total_cost >= costs["alert_threshold_usd"]:
await self._send_cost_warning(
reservation.provider,
total_cost,
costs["alert_threshold_usd"],
)
return finalized
except Exception as exc:
logger.error(
"ai_generation_receipt_finalize_failed_closed",
provider=reservation.provider,
receipt_id=reservation.receipt_id,
error_type=type(exc).__name__,
)
return False
async def get_paid_run_budget_readback(
self,
*,
run_id: str,
work_item_id: str,
expected_max_cost_micro_usd: int,
) -> dict[str, Any]:
"""Return only aggregate, non-content paid-run accounting fields."""
if any(
not _CORRELATION_VALUE_PATTERN.fullmatch(str(value or ""))
for value in (run_id, work_item_id)
):
raise ValueError("paid_run_budget_correlation_invalid")
if isinstance(expected_max_cost_micro_usd, bool):
raise ValueError("paid_run_budget_expected_cap_invalid")
expected_cap = int(expected_max_cost_micro_usd)
if not 1 <= expected_cap <= _PAID_CANARY_MAX_COST_MICRO_USD:
raise ValueError("paid_run_budget_expected_cap_invalid")
budget_hash = self._identity_hash(run_id, work_item_id)
budget_id = f"paid-run-{budget_hash[:32]}"
redis = await self._get_redis()
if redis is None:
raise RuntimeError("redis_unavailable")
raw = await redis.hgetall(PAID_RUN_BUDGET_KEY.format(budget_id=budget_id))
def _decode(value: Any) -> str:
if isinstance(value, bytes):
return value.decode("utf-8", errors="replace")
return str(value or "")
receipt = {_decode(key): _decode(value) for key, value in raw.items()}
if not receipt:
return {
"run_budget_id": budget_id,
"found": False,
"correlation_verified": False,
"cap_verified": False,
"terminal_verified": False,
}
try:
max_cost_micro_usd = int(receipt.get("max_cost_micro_usd") or 0)
reserved_cost_micro_usd = int(receipt.get("reserved_cost_micro_usd") or 0)
accounted_cost_micro_usd = int(receipt.get("accounted_cost_micro_usd") or 0)
reservation_count = int(receipt.get("reservation_count") or 0)
finalized_count = int(receipt.get("finalized_count") or 0)
except ValueError as exc:
raise ValueError("paid_run_budget_receipt_invalid") from exc
correlation_verified = bool(
receipt.get("schema_version") == "ai_paid_run_budget_receipt_v1"
and receipt.get("run_budget_id") == budget_id
and receipt.get("run_id") == run_id
and receipt.get("work_item_id") == work_item_id
)
cap_verified = max_cost_micro_usd == expected_cap
terminal_verified = bool(
correlation_verified
and cap_verified
and reserved_cost_micro_usd == 0
and 0 <= accounted_cost_micro_usd <= max_cost_micro_usd
and reservation_count > 0
and finalized_count == reservation_count
and receipt.get("status") == "finalized"
)
return {
"schema_version": receipt.get("schema_version"),
"run_budget_id": budget_id,
"found": True,
"status": receipt.get("status"),
"max_cost_micro_usd": max_cost_micro_usd,
"reserved_cost_micro_usd": reserved_cost_micro_usd,
"accounted_cost_micro_usd": accounted_cost_micro_usd,
"accounted_cost_usd": round(
accounted_cost_micro_usd / _MICRO_USD,
8,
),
"reservation_count": reservation_count,
"finalized_count": finalized_count,
"correlation_verified": correlation_verified,
"cap_verified": cap_verified,
"terminal_verified": terminal_verified,
"raw_prompt_persisted": False,
"raw_response_persisted": False,
"secret_value_present": False,
}
async def get_generation_receipt_readback(
self,
provider: str,
receipt_id: str,
*,
trace_id: str,
run_id: str,
work_item_id: str,
) -> dict[str, Any]:
"""Read one allowlisted paid-generation receipt for verification.
Generation receipts never contain prompts, responses or credentials.
This method still returns only fields needed to prove correlation,
accounting finalization and reservation release.
"""
if provider not in {"claude", "gemini"}:
raise ValueError("paid_generation_receipt_provider_invalid")
if not _CORRELATION_VALUE_PATTERN.fullmatch(str(receipt_id or "")):
raise ValueError("paid_generation_receipt_id_invalid")
expected = {
"trace_id": trace_id,
"run_id": run_id,
"work_item_id": work_item_id,
}
if any(
not _CORRELATION_VALUE_PATTERN.fullmatch(str(value or ""))
for value in expected.values()
):
raise ValueError("paid_generation_receipt_correlation_invalid")
redis = await self._get_redis()
if redis is None:
raise RuntimeError("redis_unavailable")
raw = await redis.hgetall(
GENERATION_RECEIPT_KEY.format(
provider=provider,
receipt_id=receipt_id,
)
)
def _decode(value: Any) -> str:
if isinstance(value, bytes):
return value.decode("utf-8", errors="replace")
return str(value or "")
receipt = {_decode(key): _decode(value) for key, value in raw.items()}
if not receipt:
return {
"provider": provider,
"receipt_id": receipt_id,
"found": False,
"correlation_verified": False,
"finalized_accounted": False,
"reservation_released": False,
}
correlation_verified = all(
receipt.get(field) == value for field, value in expected.items()
)
finalized_accounted = receipt.get("terminal_state") == "finalized_accounted"
reservation_released = all(
int(receipt.get(field) or 0) == 0
for field in (
"reserved_prompt_tokens",
"reserved_completion_tokens",
"reserved_total_tokens",
"reserved_cost_micro_usd",
"run_budget_reserved_cost_micro_usd",
)
)
accounted_cost_micro_usd = int(receipt.get("accounted_cost_micro_usd") or 0)
return {
"provider": provider,
"receipt_id": receipt_id,
"found": True,
"status": receipt.get("status"),
"terminal_state": receipt.get("terminal_state"),
"success": receipt.get("success") == "1",
"model": receipt.get("model"),
"pricing_version": receipt.get("pricing_version"),
"usage_source": receipt.get("usage_source"),
"accounted_tokens": int(receipt.get("accounted_tokens") or 0),
"accounted_cost_micro_usd": accounted_cost_micro_usd,
"accounted_cost_usd": round(
accounted_cost_micro_usd / _MICRO_USD,
8,
),
"run_budget_required": receipt.get("run_budget_required") == "1",
"run_budget_id": receipt.get("run_budget_id") or None,
"run_budget_max_micro_usd": int(
receipt.get("run_budget_max_micro_usd") or 0
),
"correlation_verified": correlation_verified,
"finalized_accounted": finalized_accounted,
"reservation_released": reservation_released,
"raw_prompt_persisted": False,
"raw_response_persisted": False,
"secret_value_present": False,
}
async def check_and_increment(
self,
provider: str,
tokens: int = 0,
) -> tuple[bool, str | None]:
"""
檢查並遞增計數器
Args:
provider: AI 提供者 (gemini, claude)
tokens: 本次使用的 token 數 (事後更新用)
Returns:
tuple[bool, str | None]: (是否允許, 拒絕原因)
"""
if provider in {"claude", "gemini"}:
logger.error(
"paid_provider_legacy_rate_gate_rejected",
provider=provider,
reason="reserve_generation_required",
)
return False, "reserve_generation_required"
if provider not in RATE_LIMITS:
return True, None # 無限制的 provider (如 ollama)
limits, cost_limits = self._provider_limits(provider)
r = await self._get_redis()
today = self._get_today()
# 0. 🔴 2026-03-29 ogt: 檢查累積成本 (最高優先級)
if cost_limits:
cost_limit = cost_limits["total_cost_usd"]
total_cost_key = TOTAL_COST_KEY.format(provider=provider)
current_cost = await r.get(total_cost_key)
current_cost = float(current_cost) if current_cost else 0.0
if current_cost >= cost_limit:
logger.error(
"ai_cost_limit_exceeded_blocking",
provider=provider,
current_cost=f"${current_cost:.4f}",
limit=f"${cost_limit:.2f}",
action="AUTO_SWITCH_TO_OLLAMA",
)
# 發送告警 (只發一次)
await self._send_cost_alert(provider, current_cost, cost_limit)
return (
False,
f"🔴 成本超限! ${current_cost:.2f} >= ${cost_limit:.2f},已自動切換到 Ollama",
)
# 1. 檢查 RPM
rpm_key = RPM_KEY.format(provider=provider)
current_rpm = await r.get(rpm_key)
current_rpm = int(current_rpm) if current_rpm else 0
if current_rpm >= limits["rpm"]:
logger.warning(
"ai_rate_limit_rpm",
provider=provider,
current=current_rpm,
limit=limits["rpm"],
)
return False, f"RPM limit exceeded ({current_rpm}/{limits['rpm']})"
# 2. 檢查每日請求數
daily_req_key = DAILY_REQ_KEY.format(provider=provider, date=today)
current_daily = await r.get(daily_req_key)
current_daily = int(current_daily) if current_daily else 0
if current_daily >= limits["daily_requests"]:
logger.warning(
"ai_rate_limit_daily",
provider=provider,
current=current_daily,
limit=limits["daily_requests"],
)
return (
False,
f"Daily request limit exceeded ({current_daily}/{limits['daily_requests']})",
)
# 3. 檢查每日 Token (如果有追蹤)
daily_token_key = DAILY_TOKEN_KEY.format(provider=provider, date=today)
current_tokens = await r.get(daily_token_key)
current_tokens = int(current_tokens) if current_tokens else 0
if current_tokens >= limits["daily_tokens"]:
logger.warning(
"ai_rate_limit_tokens",
provider=provider,
current=current_tokens,
limit=limits["daily_tokens"],
)
return (
False,
f"Daily token limit exceeded ({current_tokens}/{limits['daily_tokens']})",
)
# 4. 遞增計數器
pipe = r.pipeline()
# RPM: 60 秒過期
pipe.incr(rpm_key)
pipe.expire(rpm_key, 60)
# Daily requests: 明天過期
pipe.incr(daily_req_key)
pipe.expire(daily_req_key, 86400)
# Daily tokens
if tokens > 0:
pipe.incrby(daily_token_key, tokens)
pipe.expire(daily_token_key, 86400)
await pipe.execute()
logger.debug(
"ai_rate_check_passed",
provider=provider,
rpm=current_rpm + 1,
daily=current_daily + 1,
)
return True, None
async def record_cost(self, provider: str, cost_usd: float) -> None:
"""
2026-03-29 ogt: 記錄累積成本
Args:
provider: AI 提供者
cost_usd: 本次成本 (USD)
"""
if provider not in COST_LIMITS or cost_usd <= 0:
return
r = await self._get_redis()
total_cost_key = TOTAL_COST_KEY.format(provider=provider)
# 使用 INCRBYFLOAT 原子操作
new_total = await r.incrbyfloat(total_cost_key, cost_usd)
logger.info(
"ai_cost_recorded",
provider=provider,
cost_usd=f"${cost_usd:.6f}",
total_cost=f"${new_total:.4f}",
)
# 檢查是否需要發送警告 (接近上限)
_, cost_limits = self._provider_limits(provider)
alert_threshold = cost_limits["alert_threshold_usd"]
if new_total >= alert_threshold:
await self._send_cost_warning(provider, new_total, alert_threshold)
async def _send_cost_alert(
self, provider: str, current_cost: float, limit: float
) -> None:
"""
2026-03-29 ogt: 發送成本超限告警到 Telegram (只發一次)
"""
r = await self._get_redis()
alert_sent_key = COST_ALERT_SENT_KEY.format(provider=provider)
# 檢查是否已發送
if await r.get(alert_sent_key):
return
try:
from src.services.telegram_gateway import (
_telegram_send_delivery_succeeded,
get_telegram_gateway,
)
gateway = get_telegram_gateway()
if not gateway.canonical_destination_chat_id(
product_id="awoooi",
signal_family="business_finops",
severity="P2",
):
logger.warning(
"ai_cost_alert_no_send",
provider=provider,
reason="canonical_route_unavailable",
)
return
message = (
f"🚨🚨🚨 AI 成本超限警報 🚨🚨🚨\n\n"
f"Provider: {provider.upper()}\n"
f"累積成本: ${current_cost:.2f}\n"
f"上限: ${limit:.2f}\n\n"
f"⚡ 已自動切換到 Ollama\n\n"
f"如需恢復 {provider.upper()},請執行:\n"
f"redis-cli DEL ai_rate:total_cost:{provider}"
)
result = await gateway.send_canonical_message(
product_id="awoooi",
signal_family="business_finops",
severity="P2",
text=message,
parse_mode="HTML",
)
if _telegram_send_delivery_succeeded(result):
await r.set(alert_sent_key, "1", ex=86400)
logger.warning(
"ai_cost_alert_sent",
provider=provider,
current_cost=f"${current_cost:.2f}",
limit=f"${limit:.2f}",
delivery_status=result.get("_awooop_delivery_status", "sent"),
)
else:
logger.warning(
"ai_cost_alert_no_send",
provider=provider,
delivery_status=(
result.get(
"_awooop_delivery_status",
"missing_provider_ack",
)
if isinstance(result, dict)
else "missing_provider_ack"
),
)
except Exception as e:
logger.error("ai_cost_alert_failed", error=str(e))
async def _send_cost_warning(
self, provider: str, current_cost: float, threshold: float
) -> None:
"""
2026-03-29 ogt: 發送成本接近上限警告
"""
r = await self._get_redis()
warning_key = f"{REDIS_KEY_PREFIX}cost_warning_sent:{provider}"
# 每小時只發一次警告
if await r.get(warning_key):
return
try:
from src.services.telegram_gateway import (
_telegram_send_delivery_succeeded,
get_telegram_gateway,
)
gateway = get_telegram_gateway()
if not gateway.canonical_destination_chat_id(
product_id="awoooi",
signal_family="business_finops",
severity="P2",
):
logger.warning(
"ai_cost_warning_no_send",
provider=provider,
reason="canonical_route_unavailable",
)
return
_, cost_limits = self._provider_limits(provider)
limit = cost_limits["total_cost_usd"]
remaining = limit - current_cost
message = (
f"⚠️ AI 成本警告\n\n"
f"Provider: {provider.upper()}\n"
f"累積成本曝險(含 pending reservation): ${current_cost:.2f} / ${limit:.2f}\n"
f"剩餘額度: ${remaining:.2f}\n\n"
f"接近上限,請注意監控!"
)
result = await gateway.send_canonical_message(
product_id="awoooi",
signal_family="business_finops",
severity="P2",
text=message,
parse_mode="HTML",
)
if _telegram_send_delivery_succeeded(result):
await r.set(warning_key, "1", ex=3600)
logger.warning(
"ai_cost_warning_sent",
provider=provider,
current_cost=f"${current_cost:.2f}",
threshold=f"${threshold:.2f}",
)
else:
logger.warning(
"ai_cost_warning_no_send",
provider=provider,
delivery_status=(
result.get(
"_awooop_delivery_status",
"missing_provider_ack",
)
if isinstance(result, dict)
else "missing_provider_ack"
),
)
except Exception as e:
logger.warning("ai_cost_warning_failed", error=str(e))
async def record_tokens(self, provider: str, tokens: int) -> None:
"""
記錄 Token 用量 (回應後呼叫)
Args:
provider: AI 提供者
tokens: 使用的 token 數
"""
if provider not in RATE_LIMITS or tokens <= 0:
return
r = await self._get_redis()
today = self._get_today()
daily_token_key = DAILY_TOKEN_KEY.format(provider=provider, date=today)
await r.incrby(daily_token_key, tokens)
await r.expire(daily_token_key, 86400)
logger.debug(
"ai_tokens_recorded",
provider=provider,
tokens=tokens,
)
async def get_provider_authentication_status(self, provider: str) -> dict[str, Any]:
"""Read public-safe authentication evidence; configuration is separate."""
from src.core.config import settings
configured = (
bool(getattr(settings, f"{provider.upper()}_API_KEY", ""))
if provider in {"claude", "gemini"}
else False
)
base = {
"provider": provider,
"configured": configured,
"authenticated": None,
"authentication_verified": False,
"verification_status": (
"configured_not_verified" if configured else "not_configured"
),
"verification_source": None,
"verified_at": None,
}
try:
redis = await self._get_redis()
if redis is None:
raise RuntimeError("redis_unavailable")
raw = await redis.hgetall(
PROVIDER_AUTH_STATUS_KEY.format(provider=provider)
)
except Exception as exc:
base["verification_status"] = (
"authentication_readback_unavailable"
if configured
else "not_configured"
)
base["readback_error"] = type(exc).__name__
return base
if not raw:
return base
decoded = {
(key.decode() if isinstance(key, bytes) else str(key)): (
value.decode() if isinstance(value, bytes) else str(value)
)
for key, value in raw.items()
}
authenticated_value = decoded.get("authenticated", "unknown")
authenticated: bool | None = None
if authenticated_value == "true":
authenticated = True
elif authenticated_value == "false":
authenticated = False
base.update(
authenticated=authenticated,
authentication_verified=(decoded.get("authentication_verified") == "true"),
verification_status=decoded.get(
"verification_status", "configured_not_verified"
),
verification_source=decoded.get("verification_source"),
verified_at=decoded.get("verified_at"),
)
return base
async def get_usage_stats(self, provider: str) -> dict:
"""
取得用量統計 (含成本)
Args:
provider: AI 提供者
Returns:
dict: 用量統計
"""
if provider not in RATE_LIMITS:
return {"provider": provider, "limited": False}
limits, cost_limits = self._provider_limits(provider)
r = await self._get_redis()
today = self._get_today()
rpm_key = RPM_KEY.format(provider=provider)
daily_req_key = DAILY_REQ_KEY.format(provider=provider, date=today)
daily_token_key = DAILY_TOKEN_KEY.format(provider=provider, date=today)
daily_token_reserved_key = DAILY_TOKEN_RESERVED_KEY.format(
provider=provider, date=today
)
daily_cost_key = DAILY_COST_MICRO_USD_KEY.format(provider=provider, date=today)
daily_cost_reserved_key = DAILY_COST_RESERVED_MICRO_USD_KEY.format(
provider=provider,
date=today,
)
total_cost_key = TOTAL_COST_KEY.format(provider=provider)
total_cost_reserved_key = TOTAL_COST_RESERVED_MICRO_USD_KEY.format(
provider=provider
)
current_rpm = await r.get(rpm_key)
current_daily = await r.get(daily_req_key)
current_tokens = await r.get(daily_token_key)
reserved_tokens = await r.get(daily_token_reserved_key)
current_daily_cost = await r.get(daily_cost_key)
reserved_daily_cost = await r.get(daily_cost_reserved_key)
current_cost = await r.get(total_cost_key)
reserved_total_cost = await r.get(total_cost_reserved_key)
request_count = int(current_daily) if current_daily else 0
token_count = int(current_tokens) if current_tokens else 0
token_reserved_count = int(reserved_tokens) if reserved_tokens else 0
daily_cost_usd = (
int(current_daily_cost) if current_daily_cost else 0
) / _MICRO_USD
daily_cost_reserved_usd = (
int(reserved_daily_cost) if reserved_daily_cost else 0
) / _MICRO_USD
total_cost_reserved_usd = (
int(reserved_total_cost) if reserved_total_cost else 0
) / _MICRO_USD
# 2026-03-29 ogt: 加入成本資訊
cost_info = {}
if cost_limits:
cost_limit = cost_limits
current_cost_float = float(current_cost) if current_cost else 0.0
total_cost_exposure = current_cost_float + total_cost_reserved_usd
alert_threshold = cost_limit["alert_threshold_usd"]
cost_info = {
"daily_cost_usd": {
"current": round(daily_cost_usd, 6),
"reserved": round(daily_cost_reserved_usd, 6),
"limit": cost_limit.get("daily_cost_usd"),
},
"total_cost_usd": {
"current": round(current_cost_float, 4),
"reserved": round(total_cost_reserved_usd, 6),
"limit": cost_limit["total_cost_usd"],
"remaining": round(
cost_limit["total_cost_usd"] - total_cost_exposure,
4,
),
"exposure_including_pending": round(total_cost_exposure, 6),
"alert_threshold": alert_threshold,
"alert_threshold_reached": (total_cost_exposure >= alert_threshold),
},
"cost_exceeded": (total_cost_exposure >= cost_limit["total_cost_usd"]),
"cost_alert": {
"basis": "actual_plus_pending_reservations",
"exposure_usd": round(total_cost_exposure, 6),
"threshold_usd": alert_threshold,
"threshold_reached": total_cost_exposure >= alert_threshold,
},
}
pricing_info: dict[str, Any] = {}
pricing_missing = False
if provider in {"claude", "gemini"}:
from src.core.config import settings
configured_model = str(
getattr(settings, f"{provider.upper()}_MODEL")
).strip()
pricing = get_paid_provider_pricing_policy(provider, configured_model)
pricing_missing = pricing is None
pricing_info = {
"pricing_policy": {
"model": configured_model,
"supported": not pricing_missing,
"source": pricing.source if pricing else None,
"version": pricing.version if pricing else None,
"checked_at": pricing.checked_at if pricing else None,
"input_usd_per_million": (
pricing.input_usd_per_million if pricing else None
),
"output_usd_per_million": (
pricing.output_usd_per_million if pricing else None
),
},
}
accounting_gap = (
request_count > 0 and token_count <= 0 and token_reserved_count <= 0
)
authentication = await self.get_provider_authentication_status(provider)
return {
"provider": provider,
"date": today,
"rpm": {
"current": int(current_rpm) if current_rpm else 0,
"limit": limits["rpm"],
},
"daily_requests": {
"current": request_count,
"limit": limits["daily_requests"],
},
"daily_tokens": {
"current": token_count,
"reserved": token_reserved_count,
"limit": limits["daily_tokens"],
},
"accounting": {
"complete": not accounting_gap and not pricing_missing,
"status": (
"blocked_pricing_policy_missing"
if pricing_missing
else (
"degraded_requests_without_token_receipts"
if accounting_gap
else "receipt_backed"
)
),
"zero_usage_is_healthy": request_count == 0,
},
"authentication": authentication,
**cost_info,
**pricing_info,
}
async def reset_cost(self, provider: str) -> None:
"""
2026-03-29 ogt: 重置累積成本 (統帥授權後使用)
Args:
provider: AI 提供者
"""
r = await self._get_redis()
total_cost_key = TOTAL_COST_KEY.format(provider=provider)
total_cost_reserved_key = TOTAL_COST_RESERVED_MICRO_USD_KEY.format(
provider=provider
)
alert_sent_key = COST_ALERT_SENT_KEY.format(provider=provider)
await r.delete(total_cost_key, total_cost_reserved_key, alert_sent_key)
logger.info("ai_cost_reset", provider=provider)
async def reset_limits(self, provider: str) -> None:
"""
重置限制 (緊急用)
Args:
provider: AI 提供者
"""
r = await self._get_redis()
today = self._get_today()
keys = [
RPM_KEY.format(provider=provider),
DAILY_REQ_KEY.format(provider=provider, date=today),
DAILY_TOKEN_KEY.format(provider=provider, date=today),
DAILY_TOKEN_RESERVED_KEY.format(provider=provider, date=today),
DAILY_COST_MICRO_USD_KEY.format(provider=provider, date=today),
DAILY_COST_RESERVED_MICRO_USD_KEY.format(provider=provider, date=today),
]
await r.delete(*keys)
logger.info("ai_rate_limits_reset", provider=provider)
# =============================================================================
# Singleton
# =============================================================================
_rate_limiter: AIRateLimiter | None = None
def get_ai_rate_limiter() -> AIRateLimiter:
"""取得 Rate Limiter 單例"""
global _rate_limiter
if _rate_limiter is None:
_rate_limiter = AIRateLimiter()
return _rate_limiter