Files
cmhub/apps/ai/catalog.py
T

96 lines
3.2 KiB
Python
Raw Normal View History

2026-07-04 10:14:16 +08:00
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:
2026-07-16 14:13:19 +08:00
required_capabilities = REQUIRED_CAPABILITIES.get(model_alias.operation_type)
if required_capabilities is None:
2026-07-04 10:14:16 +08:00
return False
2026-07-16 14:13:19 +08:00
return required_capabilities.issubset(model_alias.ai_model.capabilities_set())
2026-07-04 10:14:16 +08:00
def _requires_image_input(ai_model: AiModel, operation_type: str) -> bool:
2026-07-16 14:13:19 +08:00
if operation_type == ModelAlias.OperationType.VISION:
return True
2026-07-04 10:14:16 +08:00
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