Files

147 lines
6.5 KiB
Python
Raw Permalink Normal View History

"""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))