Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,200 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
@Author: 臧成龙
|
||||
@Contact: 939589097@qq.com
|
||||
@Time: 2025-12-31
|
||||
@File: dumpdata.py
|
||||
@Desc: 数据导出脚本 - 类似 Django 的 dumpdata - 使用方法: python scripts/dumpdata.py [app_name] > data.json
|
||||
"""
|
||||
"""
|
||||
数据导出脚本 - 类似 Django 的 dumpdata
|
||||
使用方法: python scripts/dumpdata.py [app_name] > data.json
|
||||
"""
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
|
||||
# 添加项目根目录到 Python 路径
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
from sqlalchemy import inspect
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from app.database import AsyncSessionLocal, Base
|
||||
|
||||
|
||||
def auto_import_models():
|
||||
"""自动导入所有模型"""
|
||||
import importlib
|
||||
project_root = Path(__file__).parent.parent
|
||||
|
||||
# 需要扫描的目录
|
||||
scan_dirs = ["zq_demo", "core", "scheduler", "online_dev", "ai_platform"]
|
||||
|
||||
for scan_dir in scan_dirs:
|
||||
scan_path = project_root / scan_dir
|
||||
if not scan_path.exists():
|
||||
continue
|
||||
|
||||
# 递归查找所有 model.py 文件
|
||||
for model_file in scan_path.rglob("*model.py"):
|
||||
# 计算模块路径
|
||||
relative_path = model_file.relative_to(project_root)
|
||||
module_path = str(relative_path.with_suffix("")).replace("/", ".").replace("\\", ".")
|
||||
|
||||
try:
|
||||
importlib.import_module(module_path)
|
||||
print(f"导入模型: {module_path}", file=sys.stderr)
|
||||
except ImportError as e:
|
||||
print(f"警告: 导入失败 {module_path}: {e}", file=sys.stderr)
|
||||
|
||||
|
||||
# 自动导入所有模型
|
||||
auto_import_models()
|
||||
|
||||
|
||||
class DateTimeEncoder(json.JSONEncoder):
|
||||
"""自定义 JSON 编码器,处理日期时间和 Decimal"""
|
||||
def default(self, obj):
|
||||
if isinstance(obj, datetime):
|
||||
return obj.isoformat()
|
||||
if isinstance(obj, Decimal):
|
||||
return float(obj)
|
||||
# 处理 date 类型
|
||||
from datetime import date
|
||||
if isinstance(obj, date):
|
||||
return obj.isoformat()
|
||||
return super().default(obj)
|
||||
|
||||
|
||||
async def dump_table(session: AsyncSession, model_class):
|
||||
"""导出单个表的数据"""
|
||||
from sqlalchemy import select
|
||||
|
||||
result = await session.execute(select(model_class))
|
||||
items = result.scalars().all()
|
||||
|
||||
table_data = []
|
||||
for item in items:
|
||||
# 获取所有列
|
||||
item_dict = {}
|
||||
for column in inspect(model_class).columns:
|
||||
value = getattr(item, column.name)
|
||||
item_dict[column.name] = value
|
||||
|
||||
table_data.append({
|
||||
"model": f"{model_class.__module__}.{model_class.__name__}",
|
||||
"pk": item.id,
|
||||
"fields": item_dict
|
||||
})
|
||||
|
||||
return table_data
|
||||
|
||||
|
||||
async def dump_all_data(app_name: str = None, exclude_tables: list = None):
|
||||
"""导出所有数据或指定应用的数据
|
||||
|
||||
Args:
|
||||
app_name: 应用名称(可选),如 core、scheduler
|
||||
exclude_tables: 排除的表名列表(可选)
|
||||
"""
|
||||
import logging
|
||||
|
||||
exclude_tables = exclude_tables or []
|
||||
|
||||
# 临时禁用 SQLAlchemy 的日志输出
|
||||
sqlalchemy_logger = logging.getLogger('sqlalchemy.engine')
|
||||
original_level = sqlalchemy_logger.level
|
||||
sqlalchemy_logger.setLevel(logging.WARNING)
|
||||
|
||||
all_data = []
|
||||
|
||||
try:
|
||||
async with AsyncSessionLocal() as session:
|
||||
# 获取所有模型
|
||||
models = []
|
||||
for mapper in Base.registry.mappers:
|
||||
model_class = mapper.class_
|
||||
|
||||
# 如果指定了 app_name,只导出该应用的模型
|
||||
if app_name:
|
||||
module_name = model_class.__module__
|
||||
if not module_name.startswith(app_name):
|
||||
continue
|
||||
|
||||
# 排除指定的表
|
||||
if model_class.__tablename__ in exclude_tables:
|
||||
print(f"跳过表: {model_class.__tablename__}", file=sys.stderr)
|
||||
continue
|
||||
|
||||
models.append(model_class)
|
||||
|
||||
# 按表名排序
|
||||
models.sort(key=lambda m: m.__tablename__)
|
||||
|
||||
# 导出每个表
|
||||
for model_class in models:
|
||||
print(f"导出表: {model_class.__tablename__}", file=sys.stderr)
|
||||
table_data = await dump_table(session, model_class)
|
||||
all_data.extend(table_data)
|
||||
print(f" - 导出 {len(table_data)} 条记录", file=sys.stderr)
|
||||
finally:
|
||||
# 恢复 SQLAlchemy 日志级别
|
||||
sqlalchemy_logger.setLevel(original_level)
|
||||
|
||||
return all_data
|
||||
|
||||
|
||||
async def main():
|
||||
"""主函数"""
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description='导出数据到 JSON 文件')
|
||||
parser.add_argument('app_name', nargs='?', help='应用名称(可选),如 core、scheduler')
|
||||
parser.add_argument('-o', '--output', help='输出文件路径(可选),不指定则输出到标准输出')
|
||||
parser.add_argument('-f', '--force', action='store_true', help='强制覆盖已存在的文件')
|
||||
parser.add_argument('-e', '--exclude', action='append', default=[],
|
||||
help='排除的表名,可多次使用,如: -e scheduler_log -e core_operation_log')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.app_name:
|
||||
print(f"导出应用: {args.app_name}", file=sys.stderr)
|
||||
else:
|
||||
print("导出所有数据", file=sys.stderr)
|
||||
|
||||
if args.exclude:
|
||||
print(f"排除表: {', '.join(args.exclude)}", file=sys.stderr)
|
||||
|
||||
data = await dump_all_data(args.app_name, exclude_tables=args.exclude)
|
||||
|
||||
# 生成 JSON 字符串
|
||||
json_str = json.dumps(data, ensure_ascii=False, indent=2, cls=DateTimeEncoder)
|
||||
|
||||
# 输出到文件或标准输出
|
||||
if args.output:
|
||||
output_path = Path(args.output)
|
||||
|
||||
# 检查文件是否存在
|
||||
if output_path.exists() and not args.force:
|
||||
print(f"\n错误: 文件已存在: {output_path}", file=sys.stderr)
|
||||
print("使用 -f 或 --force 参数强制覆盖", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
# 写入文件
|
||||
with open(output_path, 'w', encoding='utf-8') as f:
|
||||
f.write(json_str)
|
||||
|
||||
print(f"\n总计导出 {len(data)} 条记录", file=sys.stderr)
|
||||
print(f"已保存到: {output_path}", file=sys.stderr)
|
||||
else:
|
||||
# 输出到标准输出
|
||||
print(json_str)
|
||||
print(f"\n总计导出 {len(data)} 条记录", file=sys.stderr)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@@ -0,0 +1,166 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
@Author: 臧成龙
|
||||
@Contact: 939589097@qq.com
|
||||
@Time: 2025-12-31
|
||||
@File: loaddata.py
|
||||
@Desc: 数据导入脚本 - 类似 Django 的 loaddata - 使用方法: python scripts/loaddata.py data.json
|
||||
"""
|
||||
"""
|
||||
数据导入脚本 - 类似 Django 的 loaddata
|
||||
使用方法: python scripts/loaddata.py data.json
|
||||
"""
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from datetime import datetime
|
||||
from typing import Dict, Any
|
||||
|
||||
# 添加项目根目录到 Python 路径
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from app.database import AsyncSessionLocal, Base
|
||||
|
||||
|
||||
def auto_import_models():
|
||||
"""自动导入所有模型"""
|
||||
import importlib
|
||||
project_root = Path(__file__).parent.parent
|
||||
|
||||
# 需要扫描的目录
|
||||
scan_dirs = ["zq_demo", "core", "scheduler", "online_dev", "ai_platform"]
|
||||
|
||||
for scan_dir in scan_dirs:
|
||||
scan_path = project_root / scan_dir
|
||||
if not scan_path.exists():
|
||||
continue
|
||||
|
||||
# 递归查找所有 model.py 文件
|
||||
for model_file in scan_path.rglob("*model.py"):
|
||||
# 计算模块路径
|
||||
relative_path = model_file.relative_to(project_root)
|
||||
module_path = str(relative_path.with_suffix("")).replace("/", ".").replace("\\", ".")
|
||||
|
||||
try:
|
||||
importlib.import_module(module_path)
|
||||
except ImportError as e:
|
||||
print(f"警告: 导入失败 {module_path}: {e}")
|
||||
|
||||
|
||||
# 自动导入所有模型
|
||||
auto_import_models()
|
||||
|
||||
|
||||
def parse_value(value):
|
||||
"""自动解析值类型(日期/日期时间字符串 → date/datetime 对象)"""
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
# ISO 日期时间(如 2026-02-14T23:04:19.451515 或 2026-02-14 23:04:19)
|
||||
if len(value) >= 19 and value[4] == '-' and value[7] == '-':
|
||||
try:
|
||||
return datetime.fromisoformat(value)
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
# 短日期(如 2025-12-28)
|
||||
if len(value) == 10 and value[4] == '-' and value[7] == '-':
|
||||
try:
|
||||
return datetime.strptime(value, '%Y-%m-%d').date()
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
return value
|
||||
|
||||
|
||||
async def load_data(file_path: str):
|
||||
"""从 JSON 文件加载数据"""
|
||||
# 读取 JSON 文件
|
||||
with open(file_path, 'r', encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
|
||||
print(f"读取到 {len(data)} 条记录")
|
||||
|
||||
# 构建模型映射
|
||||
model_map: Dict[str, Any] = {}
|
||||
for mapper in Base.registry.mappers:
|
||||
model_class = mapper.class_
|
||||
model_key = f"{model_class.__module__}.{model_class.__name__}"
|
||||
model_map[model_key] = model_class
|
||||
|
||||
async with AsyncSessionLocal() as session:
|
||||
# 临时禁用外键约束(解决数据插入顺序导致的外键冲突)
|
||||
await session.execute(text("SET session_replication_role = 'replica'"))
|
||||
|
||||
success_count = 0
|
||||
error_count = 0
|
||||
|
||||
for item in data:
|
||||
try:
|
||||
model_name = item.get("model")
|
||||
fields = item.get("fields", {})
|
||||
|
||||
if model_name not in model_map:
|
||||
print(f"警告: 未找到模型 {model_name},跳过")
|
||||
error_count += 1
|
||||
continue
|
||||
|
||||
model_class = model_map[model_name]
|
||||
|
||||
# 自动转换日期时间字段
|
||||
for key, value in fields.items():
|
||||
fields[key] = parse_value(value)
|
||||
|
||||
# 创建实例
|
||||
instance = model_class(**fields)
|
||||
session.add(instance)
|
||||
|
||||
success_count += 1
|
||||
|
||||
# 每 100 条提交一次
|
||||
if success_count % 100 == 0:
|
||||
await session.commit()
|
||||
print(f"已导入 {success_count} 条记录...")
|
||||
|
||||
except Exception as e:
|
||||
print(f"错误: 导入记录失败 - {e}")
|
||||
print(f" 模型: {item.get('model')}")
|
||||
print(f" 数据: {item.get('fields')}")
|
||||
error_count += 1
|
||||
await session.rollback()
|
||||
|
||||
# 提交剩余的数据
|
||||
try:
|
||||
await session.commit()
|
||||
except Exception as e:
|
||||
print(f"提交失败: {e}")
|
||||
await session.rollback()
|
||||
|
||||
# 恢复外键约束
|
||||
await session.execute(text("SET session_replication_role = 'origin'"))
|
||||
await session.commit()
|
||||
|
||||
print(f"\n导入完成:")
|
||||
print(f" 成功: {success_count} 条")
|
||||
print(f" 失败: {error_count} 条")
|
||||
|
||||
|
||||
async def main():
|
||||
"""主函数"""
|
||||
if len(sys.argv) < 2:
|
||||
print("用法: python scripts/loaddata.py <json_file>")
|
||||
sys.exit(1)
|
||||
|
||||
file_path = sys.argv[1]
|
||||
|
||||
if not Path(file_path).exists():
|
||||
print(f"错误: 文件不存在 - {file_path}")
|
||||
sys.exit(1)
|
||||
|
||||
print(f"从文件导入数据: {file_path}")
|
||||
await load_data(file_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@@ -0,0 +1,73 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
Normalize legacy built-in agent records to Chinese.
|
||||
|
||||
Usage:
|
||||
python scripts/localize_builtin_agents.py
|
||||
python scripts/localize_builtin_agents.py --apply
|
||||
"""
|
||||
import argparse
|
||||
import asyncio
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.database import AsyncSessionLocal
|
||||
from ai_platform.models import Agent
|
||||
from ai_platform.services.agent_localization import normalize_builtin_agent_payload
|
||||
|
||||
|
||||
def _agent_payload(agent: Agent) -> dict:
|
||||
return {
|
||||
"name": agent.name,
|
||||
"code": agent.code,
|
||||
"description": agent.description or "",
|
||||
"persona": agent.persona or {},
|
||||
}
|
||||
|
||||
|
||||
async def localize_builtin_agents(apply: bool) -> None:
|
||||
async with AsyncSessionLocal() as db:
|
||||
result = await db.execute(
|
||||
select(Agent).where(
|
||||
Agent.is_deleted == False,
|
||||
)
|
||||
)
|
||||
agents = result.scalars().all()
|
||||
|
||||
changed = []
|
||||
for agent in agents:
|
||||
original = _agent_payload(agent)
|
||||
normalized = normalize_builtin_agent_payload(original)
|
||||
if normalized == original:
|
||||
continue
|
||||
|
||||
changed.append((agent, original, normalized))
|
||||
print(f"{agent.code}: {original['name']} -> {normalized['name']}")
|
||||
|
||||
if apply:
|
||||
agent.name = normalized["name"]
|
||||
agent.description = normalized["description"]
|
||||
agent.persona = normalized.get("persona") or {}
|
||||
|
||||
if apply:
|
||||
await db.commit()
|
||||
print(f"已更新 {len(changed)} 条内置智能体记录")
|
||||
else:
|
||||
await db.rollback()
|
||||
print(f"dry-run: 将更新 {len(changed)} 条记录,确认后加 --apply 执行")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--apply", action="store_true", help="write changes to database")
|
||||
args = parser.parse_args()
|
||||
asyncio.run(localize_builtin_agents(args.apply))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,279 @@
|
||||
#!/usr/bin/env python
|
||||
"""
|
||||
Seed or rollback the Multica organization collaboration workflow and agents.
|
||||
|
||||
Usage:
|
||||
python scripts/seed_multica_org_agents.py --dry-run
|
||||
python scripts/seed_multica_org_agents.py --apply
|
||||
python scripts/seed_multica_org_agents.py --rollback
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
||||
FIXTURE_PATH = PROJECT_ROOT / "ai_platform" / "fixtures" / "multica_org_agents.json"
|
||||
|
||||
|
||||
def load_fixture() -> Dict[str, Any]:
|
||||
with FIXTURE_PATH.open("r", encoding="utf-8") as fixture_file:
|
||||
return json.load(fixture_file)
|
||||
|
||||
|
||||
def build_workflow_payload(fixture: Dict[str, Any]) -> Dict[str, Any]:
|
||||
workflow = deepcopy(fixture["workflow"])
|
||||
definition = workflow["definition"]
|
||||
workflow["published_definition"] = deepcopy(definition)
|
||||
workflow["published_at"] = datetime.utcnow()
|
||||
return workflow
|
||||
|
||||
|
||||
def build_agent_payload(agent_fixture: Dict[str, Any], workflow_id: str | None) -> Dict[str, Any]:
|
||||
payload = deepcopy(agent_fixture)
|
||||
workflow_code = payload.pop("workflow_code", None)
|
||||
if workflow_code:
|
||||
payload["workflow_id"] = workflow_id
|
||||
return payload
|
||||
|
||||
|
||||
def validate_fixture(fixture: Dict[str, Any]) -> Dict[str, int]:
|
||||
agents = fixture.get("agents", [])
|
||||
definition = fixture.get("workflow", {}).get("definition", {})
|
||||
nodes = definition.get("nodes", [])
|
||||
edges = definition.get("edges", [])
|
||||
agent_codes = {agent.get("code") for agent in agents}
|
||||
workflow_agent_codes = {
|
||||
node.get("agent_code")
|
||||
for node in nodes
|
||||
if isinstance(node, dict) and node.get("agent_code")
|
||||
}
|
||||
node_types = {node.get("type") for node in nodes}
|
||||
|
||||
required_agent_codes = {
|
||||
"multica_product_manager",
|
||||
"business_requirements_analyst",
|
||||
"system_architect",
|
||||
"frontend_engineer",
|
||||
"backend_engineer",
|
||||
"qa_engineer",
|
||||
"project_manager",
|
||||
}
|
||||
required_persona_fields = {
|
||||
"role",
|
||||
"skills",
|
||||
"constraints",
|
||||
"background",
|
||||
"examples",
|
||||
}
|
||||
required_node_types = {"start", "end", "condition", "template", "parallel", "merge"}
|
||||
|
||||
missing_agents = sorted(required_agent_codes - agent_codes)
|
||||
extra_agents = sorted(agent_codes - required_agent_codes)
|
||||
missing_workflow_agents = sorted(workflow_agent_codes - agent_codes)
|
||||
missing_persona_fields = {
|
||||
agent.get("code"): sorted(required_persona_fields - set((agent.get("persona") or {}).keys()))
|
||||
for agent in agents
|
||||
if required_persona_fields - set((agent.get("persona") or {}).keys())
|
||||
}
|
||||
missing_node_types = sorted(required_node_types - node_types)
|
||||
project_manager = next(
|
||||
(agent for agent in agents if agent.get("code") == "project_manager"),
|
||||
None,
|
||||
)
|
||||
project_manager_workflow = (project_manager or {}).get("workflow_code")
|
||||
if (
|
||||
missing_agents
|
||||
or extra_agents
|
||||
or missing_workflow_agents
|
||||
or missing_persona_fields
|
||||
or missing_node_types
|
||||
or project_manager_workflow != "multica_org_collaboration_flow"
|
||||
):
|
||||
raise ValueError(
|
||||
f"Fixture validation failed: missing_agents={missing_agents}, "
|
||||
f"extra_agents={extra_agents}, "
|
||||
f"missing_workflow_agents={missing_workflow_agents}, "
|
||||
f"missing_persona_fields={missing_persona_fields}, "
|
||||
f"missing_node_types={missing_node_types}, "
|
||||
f"project_manager_workflow={project_manager_workflow}"
|
||||
)
|
||||
|
||||
return {
|
||||
"agents": len(agents),
|
||||
"workflow_nodes": len(nodes),
|
||||
"workflow_edges": len(edges),
|
||||
"workflow_agent_refs": len(workflow_agent_codes),
|
||||
}
|
||||
|
||||
|
||||
async def upsert_workflow(session, workflow_payload: Dict[str, Any]) -> AIWorkflow:
|
||||
from sqlalchemy import select
|
||||
|
||||
from ai_platform.models.workflow import AIWorkflow, AIWorkflowVersion
|
||||
|
||||
code = workflow_payload["code"]
|
||||
result = await session.execute(select(AIWorkflow).where(AIWorkflow.code == code))
|
||||
workflow = result.scalar_one_or_none()
|
||||
|
||||
if workflow:
|
||||
for key, value in workflow_payload.items():
|
||||
setattr(workflow, key, value)
|
||||
workflow.is_deleted = False
|
||||
else:
|
||||
workflow = AIWorkflow(**workflow_payload)
|
||||
session.add(workflow)
|
||||
|
||||
await session.flush()
|
||||
|
||||
result = await session.execute(
|
||||
select(AIWorkflowVersion).where(
|
||||
AIWorkflowVersion.workflow_id == workflow.id,
|
||||
AIWorkflowVersion.version == workflow.version,
|
||||
)
|
||||
)
|
||||
version = result.scalar_one_or_none()
|
||||
version_payload = {
|
||||
"workflow_id": workflow.id,
|
||||
"version": workflow.version,
|
||||
"definition": deepcopy(workflow.definition),
|
||||
"description": "Seeded Multica organization collaboration workflow",
|
||||
"published_at": workflow.published_at,
|
||||
}
|
||||
if version:
|
||||
for key, value in version_payload.items():
|
||||
setattr(version, key, value)
|
||||
version.is_deleted = False
|
||||
else:
|
||||
session.add(AIWorkflowVersion(**version_payload))
|
||||
|
||||
return workflow
|
||||
|
||||
|
||||
async def upsert_agents(session, fixture: Dict[str, Any], workflow_id: str) -> int:
|
||||
from sqlalchemy import select
|
||||
|
||||
from ai_platform.models.agent import Agent
|
||||
from app.base_model import generate_nanoid
|
||||
|
||||
count = 0
|
||||
for agent_fixture in fixture["agents"]:
|
||||
payload = build_agent_payload(agent_fixture, workflow_id)
|
||||
result = await session.execute(select(Agent).where(Agent.code == payload["code"]))
|
||||
agent = result.scalar_one_or_none()
|
||||
if agent:
|
||||
for key, value in payload.items():
|
||||
setattr(agent, key, value)
|
||||
agent.is_deleted = False
|
||||
else:
|
||||
session.add(Agent(id=generate_nanoid(), **payload))
|
||||
count += 1
|
||||
return count
|
||||
|
||||
|
||||
async def apply_seed(dry_run: bool) -> Dict[str, Any]:
|
||||
fixture = load_fixture()
|
||||
counts = validate_fixture(fixture)
|
||||
workflow_payload = build_workflow_payload(fixture)
|
||||
|
||||
if dry_run:
|
||||
return {
|
||||
"action": "dry-run",
|
||||
"workflow_code": workflow_payload["code"],
|
||||
"agent_count": len(fixture["agents"]),
|
||||
**counts,
|
||||
}
|
||||
|
||||
from app.database import AsyncSessionLocal
|
||||
|
||||
async with AsyncSessionLocal() as session:
|
||||
workflow = await upsert_workflow(session, workflow_payload)
|
||||
agent_count = await upsert_agents(session, fixture, workflow.id)
|
||||
|
||||
await session.commit()
|
||||
action = "applied"
|
||||
|
||||
return {
|
||||
"action": action,
|
||||
"workflow_code": workflow_payload["code"],
|
||||
"agent_count": agent_count,
|
||||
**counts,
|
||||
}
|
||||
|
||||
|
||||
async def rollback_seed(dry_run: bool) -> Dict[str, Any]:
|
||||
from sqlalchemy import select
|
||||
|
||||
from ai_platform.models.agent import Agent
|
||||
from ai_platform.models.workflow import AIWorkflow, AIWorkflowVersion
|
||||
from app.database import AsyncSessionLocal
|
||||
|
||||
fixture = load_fixture()
|
||||
workflow_code = fixture["workflow"]["code"]
|
||||
agent_codes = [agent["code"] for agent in fixture["agents"]]
|
||||
|
||||
async with AsyncSessionLocal() as session:
|
||||
result = await session.execute(select(AIWorkflow).where(AIWorkflow.code == workflow_code))
|
||||
workflow = result.scalar_one_or_none()
|
||||
workflow_count = 0
|
||||
version_count = 0
|
||||
if workflow:
|
||||
workflow.is_deleted = True
|
||||
workflow_count = 1
|
||||
versions = await session.execute(
|
||||
select(AIWorkflowVersion).where(AIWorkflowVersion.workflow_id == workflow.id)
|
||||
)
|
||||
for version in versions.scalars().all():
|
||||
version.is_deleted = True
|
||||
version_count += 1
|
||||
|
||||
agents = await session.execute(select(Agent).where(Agent.code.in_(agent_codes)))
|
||||
agent_count = 0
|
||||
for agent in agents.scalars().all():
|
||||
agent.is_deleted = True
|
||||
agent.workflow_id = None
|
||||
agent_count += 1
|
||||
|
||||
if dry_run:
|
||||
await session.rollback()
|
||||
action = "rollback-dry-run"
|
||||
else:
|
||||
await session.commit()
|
||||
action = "rolled-back"
|
||||
|
||||
return {
|
||||
"action": action,
|
||||
"workflow_code": workflow_code,
|
||||
"workflow_count": workflow_count,
|
||||
"workflow_version_count": version_count,
|
||||
"agent_count": agent_count,
|
||||
}
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Seed Multica organization collaboration agents")
|
||||
mode = parser.add_mutually_exclusive_group(required=True)
|
||||
mode.add_argument("--dry-run", action="store_true", help="Validate and simulate the seed without committing")
|
||||
mode.add_argument("--apply", action="store_true", help="Apply the seed to the configured database")
|
||||
mode.add_argument("--rollback", action="store_true", help="Soft-delete seeded workflow, versions, and agents")
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.rollback:
|
||||
result = await rollback_seed(dry_run=False)
|
||||
else:
|
||||
result = await apply_seed(dry_run=args.dry_run)
|
||||
|
||||
print(json.dumps(result, ensure_ascii=False, sort_keys=True))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
Reference in New Issue
Block a user