94 lines
3.1 KiB
Python
94 lines
3.1 KiB
Python
from __future__ import annotations
|
|
|
|
from collections import defaultdict
|
|
|
|
from apps.ai.aliases import REQUIRED_CAPABILITIES
|
|
from apps.ai.models import AiModel, ModelAlias
|
|
from apps.ai.providers import AiProviderConfigError, resolve_api_type
|
|
from apps.billing.models import PricingRule
|
|
|
|
|
|
PRICING_STATUS_PRICED = "priced"
|
|
PRICING_STATUS_UNPRICED = "unpriced"
|
|
PUBLIC_DEFAULT_RESOLUTION = "default"
|
|
|
|
|
|
def _has_required_capability(model_alias: ModelAlias) -> bool:
|
|
required_capability = REQUIRED_CAPABILITIES.get(model_alias.operation_type)
|
|
if required_capability is None:
|
|
return False
|
|
return required_capability in model_alias.ai_model.capabilities_set()
|
|
|
|
|
|
def _requires_image_input(ai_model: AiModel, operation_type: str) -> bool:
|
|
if operation_type != ModelAlias.OperationType.IMAGE:
|
|
return False
|
|
try:
|
|
resolved_api_type = resolve_api_type(ai_model.api_type, ai_model.url)
|
|
except AiProviderConfigError:
|
|
return False
|
|
return resolved_api_type == AiModel.ApiType.IMAGES_EDITS
|
|
|
|
|
|
def _pricing_rule_to_public_price(rule: PricingRule) -> dict:
|
|
return {
|
|
"resolution": rule.resolution or PUBLIC_DEFAULT_RESOLUTION,
|
|
"points_cost": rule.points_cost,
|
|
}
|
|
|
|
|
|
def _sort_public_prices(prices: list[dict]) -> list[dict]:
|
|
return sorted(
|
|
prices,
|
|
key=lambda price: (
|
|
price["resolution"] != PUBLIC_DEFAULT_RESOLUTION,
|
|
price["resolution"],
|
|
),
|
|
)
|
|
|
|
|
|
def get_public_model_catalog() -> list[dict]:
|
|
aliases = [
|
|
model_alias
|
|
for model_alias in ModelAlias.objects.select_related("ai_model")
|
|
.filter(is_active=True, ai_model__is_active=True)
|
|
.order_by("operation_type", "alias", "id")
|
|
if _has_required_capability(model_alias)
|
|
]
|
|
if not aliases:
|
|
return []
|
|
|
|
alias_keys = {(item.operation_type, item.alias) for item in aliases}
|
|
prices_by_alias = defaultdict(list)
|
|
pricing_rules = PricingRule.objects.filter(
|
|
is_active=True,
|
|
operation_type__in={operation for operation, _alias in alias_keys},
|
|
alias__in={alias for _operation, alias in alias_keys},
|
|
).order_by("operation_type", "alias", "resolution", "id")
|
|
|
|
for rule in pricing_rules:
|
|
key = (rule.operation_type, rule.alias)
|
|
if key in alias_keys:
|
|
prices_by_alias[key].append(_pricing_rule_to_public_price(rule))
|
|
|
|
catalog = []
|
|
for model_alias in aliases:
|
|
key = (model_alias.operation_type, model_alias.alias)
|
|
prices = _sort_public_prices(prices_by_alias.get(key, []))
|
|
catalog.append(
|
|
{
|
|
"alias": model_alias.alias,
|
|
"operation_type": model_alias.operation_type,
|
|
"capabilities": sorted(model_alias.ai_model.capabilities_set()),
|
|
"requires_image": _requires_image_input(
|
|
model_alias.ai_model,
|
|
model_alias.operation_type,
|
|
),
|
|
"pricing_status": (
|
|
PRICING_STATUS_PRICED if prices else PRICING_STATUS_UNPRICED
|
|
),
|
|
"prices": prices,
|
|
}
|
|
)
|
|
return catalog
|