feat: add periodic practice reports
This commit is contained in:
@@ -0,0 +1,52 @@
|
||||
"""periodic reports
|
||||
|
||||
Revision ID: 0019_periodic_reports
|
||||
Revises: 0018_ai_cost_tracking
|
||||
Create Date: 2026-07-31 00:00:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0019_periodic_reports"
|
||||
down_revision = "0018_ai_cost_tracking"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"sys_periodic_report",
|
||||
sa.Column("id", sa.BigInteger().with_variant(sa.Integer(), "sqlite"), primary_key=True, autoincrement=True),
|
||||
sa.Column("user_id", sa.BigInteger().with_variant(sa.Integer(), "sqlite"), nullable=False),
|
||||
sa.Column("report_type", sa.String(length=30), nullable=False),
|
||||
sa.Column("period_start", sa.DateTime(), nullable=False),
|
||||
sa.Column("period_end", sa.DateTime(), nullable=False),
|
||||
sa.Column("title", sa.String(length=160), nullable=False),
|
||||
sa.Column("content", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("source_summary_ids", sa.Text(), nullable=True),
|
||||
sa.Column("source_topic_ids", sa.Text(), nullable=True),
|
||||
sa.Column("model_name", sa.String(length=100), nullable=True),
|
||||
sa.Column("status", sa.String(length=20), nullable=False, server_default="success"),
|
||||
sa.Column("error_message", sa.Text(), nullable=True),
|
||||
sa.Column("generated_by", sa.String(length=30), nullable=False, server_default="manual"),
|
||||
sa.Column("generated_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["sys_user.id"]),
|
||||
sa.UniqueConstraint("user_id", "report_type", "period_start", "period_end", name="uq_sys_periodic_report_period"),
|
||||
)
|
||||
op.create_index("ix_sys_periodic_report_user_id", "sys_periodic_report", ["user_id"])
|
||||
op.create_index("ix_sys_periodic_report_status", "sys_periodic_report", ["status"])
|
||||
op.create_index("ix_sys_periodic_report_user_type_created", "sys_periodic_report", ["user_id", "report_type", "created_at"])
|
||||
op.create_index("ix_sys_periodic_report_status_created", "sys_periodic_report", ["status", "created_at"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_sys_periodic_report_status_created", table_name="sys_periodic_report")
|
||||
op.drop_index("ix_sys_periodic_report_user_type_created", table_name="sys_periodic_report")
|
||||
op.drop_index("ix_sys_periodic_report_status", table_name="sys_periodic_report")
|
||||
op.drop_index("ix_sys_periodic_report_user_id", table_name="sys_periodic_report")
|
||||
op.drop_table("sys_periodic_report")
|
||||
@@ -7,6 +7,7 @@ from io import BytesIO
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile, status
|
||||
from fastapi.responses import StreamingResponse
|
||||
from openpyxl import Workbook, load_workbook
|
||||
from pydantic import BaseModel
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy import extract, func, select
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -26,6 +27,7 @@ from app.services.admin_service import OperationLogService
|
||||
from app.services.entitlement_service import EntitlementService, entitlement_dict, view_from_plan
|
||||
from app.services.growth_profile_service import GrowthProfileService, growth_profile_dict, topic_dict, topic_summary_dict
|
||||
from app.services.help_card_service import help_card_dict
|
||||
from app.services.periodic_report_service import PeriodicReportService, periodic_report_dict
|
||||
from app.services.share_draft_service import share_draft_dict
|
||||
from app.api.pagination import page_result
|
||||
|
||||
@@ -34,6 +36,12 @@ router = APIRouter()
|
||||
STUDENT_TEMPLATE_HEADERS = ["手机号", "姓名", "昵称", "每日聊天额度", "状态", "有效期"]
|
||||
|
||||
|
||||
class AdminGenerateReportRequest(BaseModel):
|
||||
reportType: str
|
||||
periodStart: datetime | None = None
|
||||
periodEnd: datetime | None = None
|
||||
|
||||
|
||||
@router.get("/user/list")
|
||||
def list_users(
|
||||
keyword: str = Query(default=""),
|
||||
@@ -259,10 +267,49 @@ def user_operation_detail(
|
||||
"recentTopics": _recent_topics(db, user.id),
|
||||
"recentHelpCards": [help_card_dict(item) for item in _recent_help_cards(db, user.id)],
|
||||
"recentShareDrafts": [share_draft_dict(item) for item in _recent_share_drafts(db, user.id)],
|
||||
"recentReports": [periodic_report_dict(item) for item in PeriodicReportService.list_user_reports(db, user_id=user.id, limit=10)],
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.get("/user/{user_id}/reports")
|
||||
def user_reports(
|
||||
user_id: int,
|
||||
limit: int = Query(default=20, ge=1, le=50),
|
||||
db: Session = Depends(get_db),
|
||||
current_admin: Admin = Depends(get_current_admin),
|
||||
) -> dict:
|
||||
user = _get_user(db, user_id)
|
||||
reports = PeriodicReportService.list_user_reports(db, user_id=user.id, limit=limit)
|
||||
return api_success([periodic_report_dict(item) for item in reports])
|
||||
|
||||
|
||||
@router.post("/user/{user_id}/reports/generate")
|
||||
def generate_user_report(
|
||||
user_id: int,
|
||||
payload: AdminGenerateReportRequest,
|
||||
db: Session = Depends(get_db),
|
||||
current_admin: Admin = Depends(get_current_admin),
|
||||
) -> dict:
|
||||
user = _get_user(db, user_id)
|
||||
if payload.reportType not in ("weekly", "monthly", "stage"):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="不支持的报告类型")
|
||||
try:
|
||||
report = PeriodicReportService.generate_for_user(
|
||||
db,
|
||||
user=user,
|
||||
report_type=payload.reportType, # type: ignore[arg-type]
|
||||
period_start=payload.periodStart,
|
||||
period_end=payload.periodEnd,
|
||||
generated_by=f"admin:{current_admin.id}",
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
|
||||
OperationLogService.write(db, admin_id=current_admin.id, module="user_report", action="generate", target_id=report.id)
|
||||
db.commit()
|
||||
return api_success(periodic_report_dict(report))
|
||||
|
||||
|
||||
@router.put("/user/{user_id}")
|
||||
def update_user(
|
||||
user_id: int,
|
||||
|
||||
@@ -10,6 +10,7 @@ from app.models.user import User
|
||||
from app.schemas.user import UserProfile
|
||||
from app.services.entitlement_service import EntitlementService, entitlement_dict
|
||||
from app.services.growth_profile_service import GrowthProfileService, growth_profile_dict
|
||||
from app.services.periodic_report_service import PeriodicReportService, periodic_report_dict
|
||||
from app.services.topic_session_service import TopicSessionService
|
||||
|
||||
router = APIRouter()
|
||||
@@ -53,3 +54,13 @@ def growth_profile(
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.get("/periodic-report/list")
|
||||
def periodic_reports(
|
||||
limit: int = 20,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
) -> dict:
|
||||
reports = PeriodicReportService.list_user_reports(db, user_id=current_user.id, limit=max(1, min(limit, 50)))
|
||||
return api_success([periodic_report_dict(item) for item in reports])
|
||||
|
||||
@@ -3,7 +3,7 @@ from app.models.ai_config import ModelConfig, Prompt, SystemConfig
|
||||
from app.models.base import Base
|
||||
from app.models.chat import ChatMessage, ChatSession, TopicSession
|
||||
from app.models.entitlement import EntitlementPlan, UserEntitlement, UserEntitlementLog
|
||||
from app.models.growth import GrowthProfileRevision, ShareDraft, TeacherHelpCard, TopicSummary, UserGrowthProfile
|
||||
from app.models.growth import GrowthProfileRevision, PeriodicReport, ShareDraft, TeacherHelpCard, TopicSummary, UserGrowthProfile
|
||||
from app.models.knowledge import (
|
||||
HumanAttentionHistory,
|
||||
HumanAttentionRecord,
|
||||
@@ -48,6 +48,7 @@ __all__ = [
|
||||
"HumanAttentionRecord",
|
||||
"ModelConfig",
|
||||
"OperationLog",
|
||||
"PeriodicReport",
|
||||
"LogRetentionPolicy",
|
||||
"StorageSnapshot",
|
||||
"TopicSession",
|
||||
|
||||
@@ -100,3 +100,29 @@ class ShareDraft(Base):
|
||||
copied: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
copied_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), nullable=False)
|
||||
|
||||
|
||||
class PeriodicReport(Base):
|
||||
__tablename__ = "sys_periodic_report"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("user_id", "report_type", "period_start", "period_end", name="uq_sys_periodic_report_period"),
|
||||
Index("ix_sys_periodic_report_user_type_created", "user_id", "report_type", "created_at"),
|
||||
Index("ix_sys_periodic_report_status_created", "status", "created_at"),
|
||||
)
|
||||
|
||||
id: Mapped[int] = mapped_column(PRIMARY_KEY_TYPE, primary_key=True, autoincrement=True)
|
||||
user_id: Mapped[int] = mapped_column(ForeignKey("sys_user.id"), index=True, nullable=False)
|
||||
report_type: Mapped[str] = mapped_column(String(30), nullable=False)
|
||||
period_start: Mapped[datetime] = mapped_column(DateTime, nullable=False)
|
||||
period_end: Mapped[datetime] = mapped_column(DateTime, nullable=False)
|
||||
title: Mapped[str] = mapped_column(String(160), nullable=False)
|
||||
content: Mapped[str] = mapped_column(Text, default="", nullable=False)
|
||||
source_summary_ids: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
source_topic_ids: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
model_name: Mapped[str | None] = mapped_column(String(100), nullable=True)
|
||||
status: Mapped[str] = mapped_column(String(20), default="success", index=True, nullable=False)
|
||||
error_message: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
generated_by: Mapped[str] = mapped_column(String(30), default="manual", nullable=False)
|
||||
generated_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), nullable=False)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), nullable=False)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), onupdate=func.now(), nullable=False)
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Literal
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.chat import TopicSession
|
||||
from app.models.growth import PeriodicReport, TopicSummary, UserGrowthProfile
|
||||
from app.models.user import User
|
||||
from app.services.model_service import ModelClientService
|
||||
|
||||
ReportType = Literal["weekly", "monthly", "stage"]
|
||||
|
||||
REPORT_TYPE_LABELS = {
|
||||
"weekly": "每周实修小结",
|
||||
"monthly": "每月成长报告",
|
||||
"stage": "阶段成长总结",
|
||||
}
|
||||
|
||||
|
||||
class PeriodicReportService:
|
||||
@staticmethod
|
||||
def default_period(report_type: ReportType, now: datetime | None = None) -> tuple[datetime, datetime]:
|
||||
current = (now or datetime.now(UTC)).replace(tzinfo=None)
|
||||
if report_type == "weekly":
|
||||
end = current
|
||||
start = end - timedelta(days=7)
|
||||
elif report_type == "monthly":
|
||||
end = current
|
||||
start = end - timedelta(days=30)
|
||||
else:
|
||||
end = current
|
||||
start = end - timedelta(days=150)
|
||||
return start.replace(microsecond=0), end.replace(microsecond=0)
|
||||
|
||||
@staticmethod
|
||||
def list_user_reports(db: Session, *, user_id: int, limit: int = 20) -> list[PeriodicReport]:
|
||||
return list(
|
||||
db.scalars(
|
||||
select(PeriodicReport)
|
||||
.where(PeriodicReport.user_id == user_id)
|
||||
.order_by(PeriodicReport.period_end.desc(), PeriodicReport.id.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def generate_for_user(
|
||||
db: Session,
|
||||
*,
|
||||
user: User,
|
||||
report_type: ReportType,
|
||||
period_start: datetime | None = None,
|
||||
period_end: datetime | None = None,
|
||||
generated_by: str = "manual",
|
||||
) -> PeriodicReport:
|
||||
if report_type not in REPORT_TYPE_LABELS:
|
||||
raise ValueError("不支持的报告类型")
|
||||
if period_start is None or period_end is None:
|
||||
default_start, default_end = PeriodicReportService.default_period(report_type)
|
||||
period_start = period_start or default_start
|
||||
period_end = period_end or default_end
|
||||
period_start = period_start.replace(tzinfo=None, microsecond=0)
|
||||
period_end = period_end.replace(tzinfo=None, microsecond=0)
|
||||
if period_start >= period_end:
|
||||
raise ValueError("报告开始时间必须早于结束时间")
|
||||
|
||||
report = _get_or_create_report(db, user=user, report_type=report_type, period_start=period_start, period_end=period_end)
|
||||
summaries = _period_summaries(db, user_id=user.id, period_start=period_start, period_end=period_end)
|
||||
profile = db.scalar(select(UserGrowthProfile).where(UserGrowthProfile.user_id == user.id))
|
||||
topic_ids = sorted({int(item.topic_session_id) for item in summaries})
|
||||
summary_ids = [int(item.id) for item in summaries]
|
||||
report.source_topic_ids = json.dumps(topic_ids, ensure_ascii=False)
|
||||
report.source_summary_ids = json.dumps(summary_ids, ensure_ascii=False)
|
||||
report.generated_by = generated_by
|
||||
report.generated_at = _now()
|
||||
|
||||
if not summaries:
|
||||
report.status = "empty"
|
||||
report.error_message = None
|
||||
report.content = _empty_report_content(report_type=report_type, period_start=period_start, period_end=period_end)
|
||||
db.add(report)
|
||||
db.commit()
|
||||
db.refresh(report)
|
||||
return report
|
||||
|
||||
try:
|
||||
prompt = _report_prompt(user=user, report_type=report_type, period_start=period_start, period_end=period_end, summaries=summaries, profile=profile)
|
||||
report.content = ModelClientService.generate_text_or_raise(db, prompt).strip() or _fallback_report(summaries)
|
||||
model = ModelClientService._get_enabled_model(db)
|
||||
report.model_name = model.model_name if model else None
|
||||
report.status = "success"
|
||||
report.error_message = None
|
||||
except Exception as exc:
|
||||
report.status = "failed"
|
||||
report.error_message = str(exc)[:2000]
|
||||
report.content = _fallback_report(summaries)
|
||||
db.add(report)
|
||||
db.commit()
|
||||
db.refresh(report)
|
||||
return report
|
||||
|
||||
|
||||
def periodic_report_dict(report: PeriodicReport) -> dict:
|
||||
return {
|
||||
"id": report.id,
|
||||
"userId": report.user_id,
|
||||
"reportType": report.report_type,
|
||||
"reportTypeLabel": REPORT_TYPE_LABELS.get(report.report_type, report.report_type),
|
||||
"periodStart": report.period_start,
|
||||
"periodEnd": report.period_end,
|
||||
"title": report.title,
|
||||
"content": report.content,
|
||||
"sourceSummaryIds": _parse_json_list(report.source_summary_ids),
|
||||
"sourceTopicIds": _parse_json_list(report.source_topic_ids),
|
||||
"modelName": report.model_name,
|
||||
"status": report.status,
|
||||
"errorMessage": report.error_message,
|
||||
"generatedBy": report.generated_by,
|
||||
"generatedAt": report.generated_at,
|
||||
"createdAt": report.created_at,
|
||||
"updatedAt": report.updated_at,
|
||||
}
|
||||
|
||||
|
||||
def _get_or_create_report(
|
||||
db: Session,
|
||||
*,
|
||||
user: User,
|
||||
report_type: ReportType,
|
||||
period_start: datetime,
|
||||
period_end: datetime,
|
||||
) -> PeriodicReport:
|
||||
report = db.scalar(
|
||||
select(PeriodicReport).where(
|
||||
PeriodicReport.user_id == user.id,
|
||||
PeriodicReport.report_type == report_type,
|
||||
PeriodicReport.period_start == period_start,
|
||||
PeriodicReport.period_end == period_end,
|
||||
)
|
||||
)
|
||||
title = f"{REPORT_TYPE_LABELS[report_type]}({period_start:%Y-%m-%d} 至 {period_end:%Y-%m-%d})"
|
||||
if report is None:
|
||||
report = PeriodicReport(
|
||||
user_id=user.id,
|
||||
report_type=report_type,
|
||||
period_start=period_start,
|
||||
period_end=period_end,
|
||||
title=title,
|
||||
)
|
||||
else:
|
||||
report.title = title
|
||||
return report
|
||||
|
||||
|
||||
def _period_summaries(db: Session, *, user_id: int, period_start: datetime, period_end: datetime) -> list[TopicSummary]:
|
||||
return list(
|
||||
db.scalars(
|
||||
select(TopicSummary)
|
||||
.where(
|
||||
TopicSummary.user_id == user_id,
|
||||
TopicSummary.status == "success",
|
||||
TopicSummary.generated_at >= period_start,
|
||||
TopicSummary.generated_at <= period_end,
|
||||
)
|
||||
.order_by(TopicSummary.generated_at.asc(), TopicSummary.id.asc())
|
||||
.limit(200)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _report_prompt(
|
||||
*,
|
||||
user: User,
|
||||
report_type: ReportType,
|
||||
period_start: datetime,
|
||||
period_end: datetime,
|
||||
summaries: list[TopicSummary],
|
||||
profile: UserGrowthProfile | None,
|
||||
) -> str:
|
||||
summary_text = "\n\n".join(
|
||||
(
|
||||
f"主题摘要 {index}\n"
|
||||
f"摘要:{item.summary}\n"
|
||||
f"主要事件:{item.main_events or '无'}\n"
|
||||
f"情绪:{item.emotions or '无'}\n"
|
||||
f"身体感受:{item.body_feelings or '无'}\n"
|
||||
f"推荐功课:{item.recommended_homework or '无'}\n"
|
||||
f"看见/变化:{item.insights or '无'}\n"
|
||||
f"下一步观察:{item.next_observation or '无'}"
|
||||
)
|
||||
for index, item in enumerate(summaries, start=1)
|
||||
)
|
||||
return (
|
||||
f"请为学员“{user.name or user.nickname or user.phone}”生成一份{REPORT_TYPE_LABELS[report_type]}。\n"
|
||||
f"周期:{period_start:%Y-%m-%d %H:%M} 至 {period_end:%Y-%m-%d %H:%M}。\n\n"
|
||||
"产品定位:这是大本营千问千答的实修陪伴报告,不写成医疗诊断、心理咨询结论或营销文。\n"
|
||||
"表达方向:回到当下、回到自身、觉察情绪和身体感受,如实释放;建议适度,不要给过多术层面的复杂方案。\n"
|
||||
"请使用 Markdown 输出,结构包含:本周期主要议题、做过或被建议的功课、反复出现的情绪/身体模式、已有变化、下一步观察方向、可以带给老师确认的问题。\n"
|
||||
"如果证据不足,请明确说“本周期沉淀记录较少”,不要编造。\n\n"
|
||||
f"长期成长档案:\n{profile.profile_text if profile else '暂无'}\n\n"
|
||||
f"本周期主题摘要:\n{summary_text}"
|
||||
)
|
||||
|
||||
|
||||
def _fallback_report(summaries: list[TopicSummary]) -> str:
|
||||
lines = ["## 本周期实修小结", "", "模型生成失败时,系统先根据已沉淀主题摘要生成基础版本:"]
|
||||
for item in summaries[:20]:
|
||||
lines.extend(["", f"- {item.summary}"])
|
||||
if item.recommended_homework:
|
||||
lines.append(f" - 建议功课:{item.recommended_homework}")
|
||||
if item.next_observation:
|
||||
lines.append(f" - 下一步观察:{item.next_observation}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _empty_report_content(*, report_type: ReportType, period_start: datetime, period_end: datetime) -> str:
|
||||
return (
|
||||
f"## {REPORT_TYPE_LABELS[report_type]}\n\n"
|
||||
f"周期:{period_start:%Y-%m-%d} 至 {period_end:%Y-%m-%d}\n\n"
|
||||
"本周期还没有可用于生成报告的主题沉淀。可以在完成一次主题对话后,先点击“沉淀本主题”,再生成报告。"
|
||||
)
|
||||
|
||||
|
||||
def _parse_json_list(raw: str | None) -> list[int]:
|
||||
if not raw:
|
||||
return []
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
return []
|
||||
if not isinstance(data, list):
|
||||
return []
|
||||
return [int(item) for item in data if isinstance(item, (int, str)) and str(item).isdigit()]
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(UTC).replace(tzinfo=None)
|
||||
@@ -0,0 +1,68 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from app.models import Base
|
||||
from app.models.ai_config import SystemConfig
|
||||
from app.models.chat import ChatSession, TopicSession
|
||||
from app.models.growth import PeriodicReport, TopicSummary, UserGrowthProfile
|
||||
from app.models.user import User
|
||||
from app.services.periodic_report_service import PeriodicReportService, periodic_report_dict
|
||||
|
||||
|
||||
def _db() -> Session:
|
||||
engine = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool)
|
||||
Base.metadata.create_all(engine)
|
||||
return Session(engine)
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(UTC).replace(tzinfo=None)
|
||||
|
||||
|
||||
def test_generate_periodic_report_from_topic_summaries():
|
||||
with _db() as db:
|
||||
now = _now()
|
||||
db.add(SystemConfig(config_key="mock_model_enabled", config_value="true"))
|
||||
user = User(id=1, phone="13800000001", name="学员", daily_chat_limit=100, daily_chat_used=0)
|
||||
session = ChatSession(id=1, user_id=1, title="表达障碍", message_count=2, last_message_at=now, is_deleted=0)
|
||||
topic = TopicSession(id=1, user_id=1, chat_session_id=1, title="表达障碍", core_question="我不敢表达", status="completed")
|
||||
summary = TopicSummary(
|
||||
id=1,
|
||||
user_id=1,
|
||||
topic_session_id=1,
|
||||
summary="本周反复看到表达时身体发紧。",
|
||||
emotions="害怕、紧张",
|
||||
body_feelings="喉咙紧、胸口堵",
|
||||
recommended_homework="表达障碍练习",
|
||||
next_observation="先观察身体反应",
|
||||
generated_at=now - timedelta(days=1),
|
||||
)
|
||||
db.add_all([user, session, topic, summary, UserGrowthProfile(user_id=1, profile_text="用户常在表达前身体发紧。")])
|
||||
db.commit()
|
||||
|
||||
report = PeriodicReportService.generate_for_user(db, user=user, report_type="weekly", period_start=now - timedelta(days=7), period_end=now)
|
||||
|
||||
assert report.status == "success"
|
||||
assert report.content
|
||||
data = periodic_report_dict(report)
|
||||
assert data["sourceSummaryIds"] == [1]
|
||||
assert data["sourceTopicIds"] == [1]
|
||||
assert db.query(PeriodicReport).filter_by(user_id=1).count() == 1
|
||||
|
||||
|
||||
def test_generate_empty_periodic_report_when_no_summaries():
|
||||
with _db() as db:
|
||||
user = User(id=1, phone="13800000001", name="学员", daily_chat_limit=100, daily_chat_used=0)
|
||||
db.add(user)
|
||||
db.commit()
|
||||
|
||||
report = PeriodicReportService.generate_for_user(db, user=user, report_type="monthly")
|
||||
|
||||
assert report.status == "empty"
|
||||
assert "还没有可用于生成报告的主题沉淀" in report.content
|
||||
assert periodic_report_dict(report)["sourceSummaryIds"] == []
|
||||
Reference in New Issue
Block a user