147 lines
6.5 KiB
Python
147 lines
6.5 KiB
Python
"""Administrative commands run inside the backend container."""
|
|
|
|
import argparse
|
|
import asyncio
|
|
import getpass
|
|
import math
|
|
|
|
from sqlalchemy import delete, update
|
|
|
|
from .config import Settings
|
|
from .db import create_database
|
|
from .models import Admin, LoginSession, MCPToken, now
|
|
from .security import password_hasher
|
|
|
|
|
|
async def reset_password():
|
|
password = getpass.getpass("New admin password (12+ characters): ")
|
|
if len(password) < 12 or password != getpass.getpass("Confirm password: "):
|
|
raise SystemExit("Password too short or confirmation does not match")
|
|
engine, sessions = create_database(Settings().database_url)
|
|
async with sessions() as db:
|
|
admin = await db.get(Admin, 1)
|
|
admin.password_hash = password_hasher.hash(password)
|
|
await db.execute(delete(LoginSession))
|
|
await db.execute(update(MCPToken).where(MCPToken.revoked_at.is_(None)).values(revoked_at=now()))
|
|
await db.commit()
|
|
await engine.dispose()
|
|
print("Admin password updated; all system sessions and MCP tokens revoked.")
|
|
|
|
|
|
async def token_command(args):
|
|
import json
|
|
|
|
from sqlalchemy import select
|
|
|
|
from .mcp_api.auth import create_token
|
|
from .models import MCPToken, now
|
|
from .research.serialization import encode_snapshot
|
|
|
|
engine, sessions = create_database(Settings().database_url)
|
|
try:
|
|
async with sessions.begin() as db:
|
|
if args.command == "mcp-token-create":
|
|
row, secret = await create_token(db, args.name, args.scope, args.days)
|
|
result = {"id": row.id, "name": row.name, "scopes": row.scopes,
|
|
"expires_at": row.expires_at, "token": secret}
|
|
elif args.command == "mcp-token-revoke":
|
|
row = await db.get(MCPToken, args.token_id)
|
|
if not row:
|
|
raise ValueError("令牌不存在")
|
|
row.revoked_at = row.revoked_at or now()
|
|
result = {"id": row.id, "revoked": True}
|
|
else:
|
|
rows = list(await db.scalars(select(MCPToken).order_by(MCPToken.created_at.desc())))
|
|
result = [{k: getattr(row, k) for k in
|
|
("id", "name", "scopes", "created_at", "expires_at", "revoked_at")} for row in rows]
|
|
# Reveal only after the transaction has committed successfully.
|
|
print(json.dumps(encode_snapshot(result), ensure_ascii=False, indent=2))
|
|
finally:
|
|
await engine.dispose()
|
|
|
|
|
|
async def catalog_sync_command(args):
|
|
"""Enqueue on the existing runner, then observe without owning the upstream session."""
|
|
import json
|
|
import time
|
|
|
|
from fastapi import HTTPException
|
|
|
|
from .business import Business
|
|
from .catalog.contracts import CatalogJobInput, Scope
|
|
from .catalog.service import Catalog
|
|
from .models import Job
|
|
|
|
engine, sessions = create_database(Settings().database_url)
|
|
try:
|
|
async with sessions.begin() as db:
|
|
if args.resume_job:
|
|
if args.region or args.universe or args.delay is not None:
|
|
raise ValueError("--resume-job 不能同时指定新范围")
|
|
job = await db.get(Job, args.resume_job)
|
|
if not job or job.kind != "catalog_full_sync":
|
|
raise ValueError("只能恢复已有全量目录任务")
|
|
result = await Business(db).retry_job(job.id)
|
|
job_id = result["id"]
|
|
else:
|
|
if not args.region or not args.universe or args.delay is None:
|
|
raise ValueError("需要 --region、--universe 和 --delay")
|
|
scope = Scope(instrument_type=args.instrument_type, region=args.region,
|
|
universe=args.universe, delay=args.delay)
|
|
job = await Catalog(db).create_job(CatalogJobInput(scope=scope), full=True)
|
|
job_id = job.id
|
|
deadline, previous = time.monotonic() + args.wait_timeout, None
|
|
while True:
|
|
async with sessions() as db:
|
|
job = await db.get(Job, job_id)
|
|
data = dict(job_id=job.id, status=job.status, processed=job.processed,
|
|
total=job.total, failed=job.failed, checkpoint=job.checkpoint, error=job.error)
|
|
current = json.dumps(data, ensure_ascii=False, sort_keys=True)
|
|
if current != previous:
|
|
print(current, flush=True)
|
|
previous = current
|
|
if job.status == "completed":
|
|
return 0
|
|
if job.status in ("failed", "completed_with_errors", "cancelled"):
|
|
return 2 if job.checkpoint.get("error_code") == "invalid_scope" else 1
|
|
if job.status in ("waiting_auth", "waiting_connection"):
|
|
return 3
|
|
if time.monotonic() >= deadline:
|
|
print(f"等待超时;后台任务 {job_id} 继续执行", flush=True)
|
|
return 4
|
|
await asyncio.sleep(min(5, max(0, deadline - time.monotonic())))
|
|
except HTTPException as exc:
|
|
print(str(exc.detail), flush=True)
|
|
return 3 if exc.status_code == 409 else 1
|
|
finally:
|
|
await engine.dispose()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
parser = argparse.ArgumentParser()
|
|
commands = parser.add_subparsers(dest="command", required=True)
|
|
commands.add_parser("reset-password")
|
|
create = commands.add_parser("mcp-token-create")
|
|
create.add_argument("--name", required=True)
|
|
create.add_argument("--scope", action="append", default=None)
|
|
create.add_argument("--days", type=int, default=90)
|
|
commands.add_parser("mcp-token-list")
|
|
revoke = commands.add_parser("mcp-token-revoke")
|
|
revoke.add_argument("token_id")
|
|
sync = commands.add_parser("catalog-sync", help="全量同步一个范围的数据集及全部字段")
|
|
sync.add_argument("--region")
|
|
sync.add_argument("--universe")
|
|
sync.add_argument("--delay", type=int, choices=range(0, 10))
|
|
sync.add_argument("--instrument-type", default="EQUITY")
|
|
sync.add_argument("--resume-job")
|
|
sync.add_argument("--wait-timeout", type=float, default=21600)
|
|
args = parser.parse_args()
|
|
if args.command == "catalog-sync" and (not math.isfinite(args.wait_timeout) or args.wait_timeout <= 0):
|
|
parser.error("--wait-timeout 必须大于 0")
|
|
try:
|
|
if args.command == "catalog-sync":
|
|
raise SystemExit(asyncio.run(catalog_sync_command(args)))
|
|
asyncio.run(reset_password() if args.command == "reset-password" else token_command(args))
|
|
except ValueError as exc:
|
|
parser.error(str(exc))
|