146 lines
5.0 KiB
Python
146 lines
5.0 KiB
Python
from __future__ import annotations
|
|
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.enums.llm_billing import LlmBillingConfigKey
|
|
from app.services.llm_billing.context import LlmBillingPolicy
|
|
from app.services.system_config_cache import get_system_config_values
|
|
|
|
_DEFAULT_HOLD_CREDITS = 5.0
|
|
_FALSE_VALUES = {"0", "false", "no", "off", "disabled"}
|
|
_HOLD_CONFIG_KEYS = {
|
|
LlmBillingConfigKey.HOLD_DEFAULT.value,
|
|
LlmBillingConfigKey.HOLD_GENERATION_RECORD_PROMPT.value,
|
|
LlmBillingConfigKey.HOLD_MODULE_IMAGE_PROMPT.value,
|
|
LlmBillingConfigKey.HOLD_MODULE_VIDEO_PROMPT.value,
|
|
LlmBillingConfigKey.HOLD_SHOT_VIDEO_ANALYSIS.value,
|
|
LlmBillingConfigKey.LEGACY_OPTIMIZE_HOLD.value,
|
|
}
|
|
|
|
|
|
def _parse_bool(value: str | None, *, default: bool = True) -> bool:
|
|
if value is None or str(value).strip() == "":
|
|
return default
|
|
return str(value).strip().lower() not in _FALSE_VALUES
|
|
|
|
|
|
def _parse_float(value: str | float | int | None) -> float | None:
|
|
try:
|
|
if value is None or str(value).strip() == "":
|
|
return None
|
|
return round(float(value), 2)
|
|
except (TypeError, ValueError):
|
|
return None
|
|
|
|
|
|
def is_llm_hold_config_key(key: str | None) -> bool:
|
|
return bool(key and key in _HOLD_CONFIG_KEYS)
|
|
|
|
|
|
async def get_llm_billing_policy(
|
|
db: AsyncSession,
|
|
*,
|
|
config_key: str | None = None,
|
|
explicit_hold_credits: float | None = None,
|
|
default: float = _DEFAULT_HOLD_CREDITS,
|
|
) -> LlmBillingPolicy:
|
|
keys = [LlmBillingConfigKey.ENABLED.value]
|
|
if config_key:
|
|
keys.append(config_key)
|
|
keys.extend(
|
|
[
|
|
LlmBillingConfigKey.HOLD_DEFAULT.value,
|
|
LlmBillingConfigKey.LEGACY_OPTIMIZE_HOLD.value,
|
|
]
|
|
)
|
|
# 去重并保持优先级;一次读取避免 enabled/scene/default 分散查询。
|
|
ordered_keys = list(dict.fromkeys(keys))
|
|
values = await get_system_config_values(db, ordered_keys)
|
|
enabled = _parse_bool(values.get(LlmBillingConfigKey.ENABLED.value), default=True)
|
|
if not enabled:
|
|
return LlmBillingPolicy(enabled=False, hold_credits=0.0, config_key=config_key)
|
|
|
|
if explicit_hold_credits is not None:
|
|
amount = _parse_float(explicit_hold_credits)
|
|
source_key = "explicit"
|
|
else:
|
|
amount = None
|
|
source_key = None
|
|
for key in ordered_keys[1:]:
|
|
parsed = _parse_float(values.get(key))
|
|
if parsed is not None:
|
|
amount = parsed
|
|
source_key = key
|
|
break
|
|
if amount is None:
|
|
amount = round(float(default), 2)
|
|
source_key = "default"
|
|
|
|
if amount is None or amount <= 0:
|
|
return LlmBillingPolicy(
|
|
enabled=True,
|
|
hold_credits=float(amount or 0),
|
|
config_key=config_key,
|
|
source_key=source_key,
|
|
valid=False,
|
|
error="启用LLM统一计费时,预扣积分必须大于0",
|
|
)
|
|
return LlmBillingPolicy(
|
|
enabled=True,
|
|
hold_credits=round(float(amount), 2),
|
|
config_key=config_key,
|
|
source_key=source_key,
|
|
)
|
|
|
|
|
|
async def get_llm_hold_credits(
|
|
db: AsyncSession,
|
|
*,
|
|
config_key: str | None = None,
|
|
default: float = _DEFAULT_HOLD_CREDITS,
|
|
) -> float:
|
|
policy = await get_llm_billing_policy(db, config_key=config_key, default=default)
|
|
return policy.hold_credits
|
|
|
|
|
|
async def is_llm_billing_enabled(db: AsyncSession) -> bool:
|
|
return (await get_llm_billing_policy(db)).enabled
|
|
|
|
|
|
async def validate_llm_system_config_value(
|
|
db: AsyncSession,
|
|
*,
|
|
key: str,
|
|
value: str,
|
|
) -> None:
|
|
"""校验后台单项更新,避免启用计费时保存零或负数预扣。"""
|
|
if key == LlmBillingConfigKey.ENABLED.value:
|
|
if not _parse_bool(value, default=True):
|
|
return
|
|
keys = [
|
|
LlmBillingConfigKey.HOLD_DEFAULT.value,
|
|
LlmBillingConfigKey.HOLD_GENERATION_RECORD_PROMPT.value,
|
|
LlmBillingConfigKey.HOLD_MODULE_IMAGE_PROMPT.value,
|
|
LlmBillingConfigKey.HOLD_MODULE_VIDEO_PROMPT.value,
|
|
LlmBillingConfigKey.HOLD_SHOT_VIDEO_ANALYSIS.value,
|
|
]
|
|
values = await get_system_config_values(db, keys, ttl_seconds=1)
|
|
invalid = [
|
|
config_name
|
|
for config_name in keys
|
|
if (raw_value := values.get(config_name)) is not None
|
|
and str(raw_value).strip() != ""
|
|
and ((parsed := _parse_float(raw_value)) is None or parsed <= 0)
|
|
]
|
|
if invalid:
|
|
raise ValueError(f"启用LLM统一计费前,请先将以下预扣配置设置为大于0:{', '.join(invalid)}")
|
|
return
|
|
|
|
if not is_llm_hold_config_key(key):
|
|
return
|
|
parsed = _parse_float(value)
|
|
enabled_values = await get_system_config_values(db, [LlmBillingConfigKey.ENABLED.value], ttl_seconds=1)
|
|
enabled = _parse_bool(enabled_values.get(LlmBillingConfigKey.ENABLED.value), default=True)
|
|
if enabled and (parsed is None or parsed <= 0):
|
|
raise ValueError("启用LLM统一计费时,预扣积分必须大于0")
|