from __future__ import annotations import argparse import asyncio from sqlalchemy import select from app.models.base import async_session from app.models.menu_config import MenuConfig from app.models.model_pricing_rule import ModelPricingRule from app.services.model_pricing.rule_service import create_rule from app.services.model_pricing.seed_data import volcengine_pricing_seed_rules from app.utils.id_gen import generate_id async def _ensure_admin_menu(db) -> bool: exists = ( await db.execute( select(MenuConfig.id) .where(MenuConfig.menu_target == "admin") .where(MenuConfig.path == "/model-pricing") .limit(1) ) ).scalar_one_or_none() if exists: return False group_id = ( await db.execute( select(MenuConfig.id) .where(MenuConfig.menu_target == "admin") .where(MenuConfig.menu_type == "group") .where(MenuConfig.label == "模型设置") .limit(1) ) ).scalar_one_or_none() if not group_id: group_id = generate_id() db.add( MenuConfig( id=group_id, path="", label="模型设置", icon="RobotOutlined", sort_order=98, is_active=True, menu_type="group", menu_target="admin", ) ) await db.flush() db.add( MenuConfig( id=generate_id(), path="/model-pricing", label="模型计价", icon="DollarOutlined", sort_order=4, is_active=True, menu_type="page", menu_target="admin", parent_id=group_id, ) ) await db.flush() return True async def run(*, commit: bool) -> None: async with async_session() as db: created = skipped = 0 menu_created = await _ensure_admin_menu(db) for payload in volcengine_pricing_seed_rules(): exists = ( await db.execute( select(ModelPricingRule.id) .where(ModelPricingRule.provider == payload["provider"]) .where(ModelPricingRule.model_name == payload["model_name"]) .where(ModelPricingRule.version_code == payload["version_code"]) .limit(1) ) ).scalar_one_or_none() if exists: skipped += 1 print(f"SKIP {payload['model_name']} {payload['version_code']} id={exists}") continue draft_payload = dict(payload) draft_payload.pop("publish_status", None) snapshot = await create_rule(db, payload=draft_payload, operator_id=None) created += 1 effective_from = snapshot["effective_from"] print( f"CREATE_DRAFT {snapshot['model_name']} {snapshot['version_code']} id={snapshot['id']} " f"effective_from={effective_from.isoformat()}" ) if commit: await db.commit() print(f"COMMIT created={created} skipped={skipped} menu_created={menu_created}") else: await db.rollback() print(f"DRY-RUN created={created} skipped={skipped} menu_created={menu_created}") def main() -> None: parser = argparse.ArgumentParser(description="初始化火山模型计价草稿(不会自动发布,需人工核价后在后台发布)") group = parser.add_mutually_exclusive_group(required=True) group.add_argument("--dry-run", action="store_true") group.add_argument("--commit", action="store_true") args = parser.parse_args() asyncio.run(run(commit=bool(args.commit))) if __name__ == "__main__": main()