116 lines
3.8 KiB
Python
116 lines
3.8 KiB
Python
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()
|