3a963a130f
个股/财务/复盘报告改为 data/user_data/reports/{username}/ 目录下,
API 端从 request.state.username 读取对应用户的报告。
Co-Authored-By: Claude <noreply@anthropic.com>
88 lines
2.8 KiB
Python
88 lines
2.8 KiB
Python
"""AI 财务分析报告持久化存储 — 按用户名隔离。
|
|
|
|
存储位置: data/user_data/reports/{username}/ai_reports.json
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import time
|
|
from pathlib import Path
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
MAX_REPORTS = 20
|
|
|
|
|
|
def _path(username: str | None = None) -> Path:
|
|
from app.config import settings
|
|
if username:
|
|
p = settings.data_dir / "user_data" / "reports" / username / "ai_reports.json"
|
|
else:
|
|
p = settings.data_dir / "user_data" / "ai_reports.json"
|
|
p.parent.mkdir(parents=True, exist_ok=True)
|
|
return p
|
|
|
|
|
|
def list_reports(username: str | None = None) -> list[dict]:
|
|
"""返回指定用户的全部报告(按 created_at 降序)。"""
|
|
p = _path(username)
|
|
if not p.exists():
|
|
return []
|
|
try:
|
|
data = json.loads(p.read_text(encoding="utf-8"))
|
|
if isinstance(data, list):
|
|
return sorted(data, key=lambda r: r.get("created_at", ""), reverse=True)
|
|
except Exception as e: # noqa: BLE001
|
|
logger.warning("ai_reports.json malformed: %s", e)
|
|
return []
|
|
|
|
|
|
def _save_all(reports: list[dict], username: str | None = None) -> None:
|
|
"""全量写入(裁剪到 MAX_REPORTS)。"""
|
|
reports.sort(key=lambda r: r.get("created_at", ""), reverse=True)
|
|
if len(reports) > MAX_REPORTS:
|
|
reports = reports[:MAX_REPORTS]
|
|
_path(username).write_text(
|
|
json.dumps(reports, indent=2, ensure_ascii=False), encoding="utf-8",
|
|
)
|
|
|
|
|
|
def save_report(report: dict, username: str | None = None) -> dict:
|
|
"""新增一条报告并持久化。返回保存后的报告(含 id / created_at)。"""
|
|
reports = list_reports(username)
|
|
if not report.get("id"):
|
|
report["id"] = f"rpt_{int(time.time() * 1000)}_{report.get('symbol', 'x')}"
|
|
if not report.get("created_at"):
|
|
report["created_at"] = _now_iso()
|
|
reports.append(report)
|
|
_save_all(reports, username)
|
|
logger.info("AI report saved by %s: %s (%s), total %d", username or "?", report.get("symbol"), report.get("id"), len(reports))
|
|
return report
|
|
|
|
|
|
def delete_report(report_id: str, username: str | None = None) -> bool:
|
|
"""删除指定用户的报告。返回是否删除成功。"""
|
|
reports = list_reports(username)
|
|
before = len(reports)
|
|
reports = [r for r in reports if r.get("id") != report_id]
|
|
if len(reports) < before:
|
|
_save_all(reports, username)
|
|
return True
|
|
return False
|
|
|
|
|
|
def clear_reports(username: str | None = None) -> int:
|
|
"""清空指定用户的全部报告。返回删除数量。"""
|
|
reports = list_reports(username)
|
|
n = len(reports)
|
|
if n > 0:
|
|
_save_all([], username)
|
|
return n
|
|
|
|
|
|
def _now_iso() -> str:
|
|
from datetime import datetime
|
|
return datetime.now().isoformat(timespec="seconds")
|