feat:新增投顾agent和nl2sqlagent

This commit is contained in:
2026-09-13 16:19:24 +08:00
parent c80c6acac0
commit 163192bf55
122 changed files with 7488 additions and 362 deletions
+62
View File
@@ -0,0 +1,62 @@
"""投顾 Agent 草稿仓储。"""
from __future__ import annotations
from sqlalchemy import select
from agent.advisor_agent.drafts import (
DRAFT_STATUS_DISCARDED,
DRAFT_STATUS_DRAFT,
)
from model.advisor_draft import AdvisorDraft
from repositories.base import BaseRepository
class AdvisorDraftRepo(BaseRepository):
model = AdvisorDraft
async def get_by_draft_id(self, draft_id: str) -> AdvisorDraft | None:
return await self.db.scalar(
select(AdvisorDraft).where(AdvisorDraft.draft_id == draft_id)
)
async def save(self, draft: AdvisorDraft) -> AdvisorDraft:
self.db.add(draft)
await self.db.commit()
await self.db.refresh(draft)
return draft
async def list_drafts(
self,
*,
advisor_id: int | None = None,
customer_id: int | None = None,
status: str | None = None,
limit: int = 20,
offset: int = 0,
) -> tuple[int, list[AdvisorDraft]]:
filters = []
if advisor_id is not None:
filters.append(AdvisorDraft.advisor_id == advisor_id)
if customer_id is not None:
filters.append(AdvisorDraft.customer_id == customer_id)
if status is not None:
filters.append(AdvisorDraft.status == status)
total = await self.count(where=filters)
items = await self.list(
where=filters,
order_by=AdvisorDraft.update_time.desc(),
limit=limit,
offset=offset,
)
return total, items
async def discard(self, draft: AdvisorDraft) -> AdvisorDraft:
if draft.status == DRAFT_STATUS_DISCARDED:
return draft
if draft.status != DRAFT_STATUS_DRAFT:
raise ValueError("草稿状态无效")
draft.status = DRAFT_STATUS_DISCARDED
await self.db.commit()
await self.db.refresh(draft)
return draft
+10 -2
View File
@@ -60,12 +60,20 @@ class AdvisorReportRepo(BaseRepository):
return (await self.db.scalar(stmt)) or 0
async def list_by_customer(
self, customer_id: int, limit: int = 100, offset: int = 0
self,
customer_id: int,
*,
advisor_id: int | None = None,
limit: int = 100,
offset: int = 0,
) -> list[AdvisorReport]:
"""某客户的历史建议报告(360 视图用)。"""
conditions = [AdvisorReport.customer_id == customer_id]
if advisor_id is not None:
conditions.append(AdvisorReport.advisor_id == advisor_id)
stmt = (
select(AdvisorReport)
.where(AdvisorReport.customer_id == customer_id)
.where(*conditions)
.order_by(AdvisorReport.id.desc())
.limit(limit)
.offset(offset)
+9
View File
@@ -10,6 +10,15 @@ from repositories.base import BaseRepository
class AdvisorVisitRecordRepo(BaseRepository):
model = AdvisorVisitRecord
async def get_by_advisor(
self, visit_id: int, advisor_id: int
) -> AdvisorVisitRecord | None:
stmt = select(AdvisorVisitRecord).where(
AdvisorVisitRecord.id == visit_id,
AdvisorVisitRecord.advisor_id == advisor_id,
)
return await self.db.scalar(stmt)
async def list_by_advisor(
self,
*,
+14 -3
View File
@@ -12,10 +12,19 @@ from repositories.base import BaseRepository
class AuditLogRepo(BaseRepository):
model = AuditLog
def _conds(self, *, user_id, module, action, keyword, start, end):
def _conds(self, *, user_id, module, action, customer_id, keyword, start, end):
conds = [AuditLog.user_id == user_id, AuditLog.module == module]
if action:
conds.append(AuditLog.action == action)
if customer_id is not None:
customer = str(customer_id)
conds.append(
or_(
AuditLog.target == customer,
AuditLog.detail.like(f'%"customer_id": {customer}%'),
AuditLog.detail.like(f'%"customer_id":{customer}%'),
)
)
if start:
conds.append(AuditLog.create_time >= start)
if end:
@@ -31,6 +40,7 @@ class AuditLogRepo(BaseRepository):
user_id: int,
module: str = "advisor",
action: str | None = None,
customer_id: int | None = None,
keyword: str | None = None,
start: datetime | None = None,
end: datetime | None = None,
@@ -39,7 +49,7 @@ class AuditLogRepo(BaseRepository):
) -> list[AuditLog]:
stmt = (
select(AuditLog)
.where(*self._conds(user_id=user_id, module=module, action=action, keyword=keyword, start=start, end=end))
.where(*self._conds(user_id=user_id, module=module, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end))
.order_by(AuditLog.id.desc())
.limit(limit)
.offset(offset)
@@ -52,6 +62,7 @@ class AuditLogRepo(BaseRepository):
user_id: int,
module: str = "advisor",
action: str | None = None,
customer_id: int | None = None,
keyword: str | None = None,
start: datetime | None = None,
end: datetime | None = None,
@@ -59,6 +70,6 @@ class AuditLogRepo(BaseRepository):
stmt = (
select(func.count())
.select_from(AuditLog)
.where(*self._conds(user_id=user_id, module=module, action=action, keyword=keyword, start=start, end=end))
.where(*self._conds(user_id=user_id, module=module, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end))
)
return (await self.db.scalar(stmt)) or 0
+18
View File
@@ -11,6 +11,7 @@ from decimal import Decimal
from sqlalchemy import func, or_, select
from common.common_const import CUSTOMER_REL_STATUS_SIGNED, CUSTOMER_REL_STATUS_UNSIGNED
from model.customer_relation import CustomerRelation
from model.fin_customer_profile import FinCustomerProfile
from model.fin_holdings import FinHoldings
@@ -24,6 +25,23 @@ _HOLDING_STATUS = "持有中"
class CustomerRelationRepo(BaseRepository):
model = CustomerRelation
async def get_active_relation(
self, *, customer_id: int, advisor_id: int
) -> CustomerRelation | None:
"""读取投顾可访问的当前关系,排除已结束关系。"""
return await self.db.scalar(
select(CustomerRelation)
.where(
CustomerRelation.customer_id == customer_id,
CustomerRelation.advisor_id == advisor_id,
CustomerRelation.status.in_(
[CUSTOMER_REL_STATUS_UNSIGNED, CUSTOMER_REL_STATUS_SIGNED]
),
)
.order_by(CustomerRelation.id.desc())
.limit(1)
)
async def get_by_customer_advisor(
self, customer_id: int, advisor_id: int
) -> CustomerRelation | None:
+31
View File
@@ -0,0 +1,31 @@
"""基金业绩指标仓储。"""
from __future__ import annotations
from sqlalchemy import select
from model.fund_performance import FundPerformance
from repositories.base import BaseRepository
class FundPerformanceRepo(BaseRepository):
model = FundPerformance
async def get_latest_for_product(self, product_id: int) -> FundPerformance | None:
stmt = (
select(FundPerformance)
.where(FundPerformance.product_id == product_id)
.order_by(
FundPerformance.calc_date.desc(),
FundPerformance.id.desc(),
)
.limit(1)
)
return (await self.db.scalars(stmt)).first()
async def list_for_product(self, product_id: int) -> list[FundPerformance]:
stmt = (
select(FundPerformance)
.where(FundPerformance.product_id == product_id)
.order_by(FundPerformance.period.asc())
)
return list((await self.db.scalars(stmt)).all())
+184
View File
@@ -0,0 +1,184 @@
"""NL2SQL 查询权限仓储。"""
from __future__ import annotations
from datetime import datetime
from sqlalchemy import delete, select
from model.nl2sql_permission import (
Nl2SqlQueryRole,
Nl2SqlRoleColumnPermission,
Nl2SqlRoleTablePermission,
Nl2SqlSensitiveField,
Nl2SqlQueryHistory,
)
from repositories.base import BaseRepository
class Nl2SqlPermissionRepo(BaseRepository):
async def list_roles(self, *, include_inactive: bool = False):
stmt = select(Nl2SqlQueryRole).order_by(Nl2SqlQueryRole.id)
if not include_inactive:
stmt = stmt.where(Nl2SqlQueryRole.status == "active")
return list((await self.db.scalars(stmt)).all())
async def get_role(self, role_id: int):
return await self.db.get(Nl2SqlQueryRole, role_id)
async def add_role(self, **kwargs):
return await self.add(Nl2SqlQueryRole(**kwargs))
async def update_role(self, role_id: int, **kwargs):
obj = await self.get_role(role_id)
if obj is None:
return None
for key, value in kwargs.items():
if value is not None:
setattr(obj, key, value)
await self.db.commit()
await self.db.refresh(obj)
return obj
async def list_role_table_permissions(self, role_id: int, *, include_inactive: bool = False):
stmt = select(Nl2SqlRoleTablePermission).where(
Nl2SqlRoleTablePermission.role_id == role_id
)
if not include_inactive:
stmt = stmt.where(Nl2SqlRoleTablePermission.status == "active")
return list((await self.db.scalars(stmt.order_by(Nl2SqlRoleTablePermission.id))).all())
async def get_table_permission(self, permission_id: int):
return await self.db.get(Nl2SqlRoleTablePermission, permission_id)
async def add_table_permission(self, **kwargs):
return await self.add(Nl2SqlRoleTablePermission(**kwargs))
async def update_table_permission(self, permission_id: int, **kwargs):
obj = await self.get_table_permission(permission_id)
if obj is None:
return None
for key, value in kwargs.items():
if value is not None:
setattr(obj, key, value)
await self.db.commit()
await self.db.refresh(obj)
return obj
async def list_role_column_permissions(self, role_id: int, *, include_inactive: bool = False):
stmt = select(Nl2SqlRoleColumnPermission).where(
Nl2SqlRoleColumnPermission.role_id == role_id
)
if not include_inactive:
stmt = stmt.where(Nl2SqlRoleColumnPermission.status == "active")
return list((await self.db.scalars(stmt.order_by(Nl2SqlRoleColumnPermission.id))).all())
async def get_column_permission(self, permission_id: int):
return await self.db.get(Nl2SqlRoleColumnPermission, permission_id)
async def add_column_permission(self, **kwargs):
return await self.add(Nl2SqlRoleColumnPermission(**kwargs))
async def update_column_permission(self, permission_id: int, **kwargs):
obj = await self.get_column_permission(permission_id)
if obj is None:
return None
for key, value in kwargs.items():
if value is not None:
setattr(obj, key, value)
await self.db.commit()
await self.db.refresh(obj)
return obj
async def list_sensitive_fields_admin(self, *, include_inactive: bool = False):
stmt = select(Nl2SqlSensitiveField).order_by(Nl2SqlSensitiveField.id)
if not include_inactive:
stmt = stmt.where(Nl2SqlSensitiveField.status == "active")
return list((await self.db.scalars(stmt)).all())
async def get_sensitive_field(self, field_id: int):
return await self.db.get(Nl2SqlSensitiveField, field_id)
async def add_sensitive_field(self, **kwargs):
return await self.add(Nl2SqlSensitiveField(**kwargs))
async def update_sensitive_field(self, field_id: int, **kwargs):
obj = await self.get_sensitive_field(field_id)
if obj is None:
return None
for key, value in kwargs.items():
if value is not None:
setattr(obj, key, value)
await self.db.commit()
await self.db.refresh(obj)
return obj
async def delete_history_before(self, before: datetime) -> int:
result = await self.db.execute(
delete(Nl2SqlQueryHistory).where(Nl2SqlQueryHistory.create_time < before)
)
await self.db.commit()
return int(result.rowcount or 0)
async def get_role_by_employee_role(self, employee_role: str):
return await self.db.scalar(
select(Nl2SqlQueryRole).where(
Nl2SqlQueryRole.employee_role == employee_role,
Nl2SqlQueryRole.status == "active",
)
)
async def list_table_permissions(self, role_id: int):
result = await self.db.scalars(
select(Nl2SqlRoleTablePermission).where(
Nl2SqlRoleTablePermission.role_id == role_id,
Nl2SqlRoleTablePermission.status == "active",
)
)
return list(result.all())
async def list_column_permissions(self, role_id: int):
result = await self.db.scalars(
select(Nl2SqlRoleColumnPermission).where(
Nl2SqlRoleColumnPermission.role_id == role_id,
Nl2SqlRoleColumnPermission.status == "active",
)
)
return list(result.all())
async def list_sensitive_fields(self):
result = await self.db.scalars(
select(Nl2SqlSensitiveField).where(Nl2SqlSensitiveField.status == "active")
)
return list(result.all())
async def list_query_history(
self,
user_id: int,
*,
limit: int = 20,
offset: int = 0,
status: str | None = None,
start_time: datetime | None = None,
end_time: datetime | None = None,
):
"""按员工自身范围分页读取查询历史。"""
stmt = select(Nl2SqlQueryHistory).where(Nl2SqlQueryHistory.user_id == user_id)
if status:
stmt = stmt.where(Nl2SqlQueryHistory.status == status)
if start_time:
stmt = stmt.where(Nl2SqlQueryHistory.create_time >= start_time)
if end_time:
stmt = stmt.where(Nl2SqlQueryHistory.create_time <= end_time)
result = await self.db.scalars(
stmt.order_by(Nl2SqlQueryHistory.create_time.desc()).limit(limit).offset(offset)
)
return list(result.all())
async def get_query_history(self, user_id: int, query_id: str):
"""按员工自身范围读取单条查询历史。"""
return await self.db.scalar(
select(Nl2SqlQueryHistory).where(
Nl2SqlQueryHistory.user_id == user_id,
Nl2SqlQueryHistory.query_id == query_id,
)
)
+12
View File
@@ -24,3 +24,15 @@ class PortfolioBenchmarkRepo(BaseRepository):
)
).all()
)
async def get_active_by_risk(self, risk_level: str) -> PortfolioBenchmark | None:
"""Return the enabled benchmark used by the Agent rebalance flow."""
return await self.db.scalar(
select(PortfolioBenchmark)
.where(
PortfolioBenchmark.status == _ACTIVE_STATUS,
PortfolioBenchmark.risk_level == risk_level,
)
.order_by(PortfolioBenchmark.id.desc())
.limit(1)
)
+17
View File
@@ -1,6 +1,8 @@
"""风评域仓储:风评记录 + 客户画像(画像主键为 customer_id)。"""
from __future__ import annotations
from datetime import date
from sqlalchemy import select
from model.fin_customer_profile import FinCustomerProfile
@@ -11,6 +13,21 @@ from repositories.base import BaseRepository
class RiskAssessmentRepo(BaseRepository):
model = FinRiskAssessment
async def get_current_by_customer(self, customer_id: int) -> FinRiskAssessment | None:
"""Return the latest currently valid risk assessment for a customer."""
return await self.db.scalar(
select(FinRiskAssessment)
.where(
FinRiskAssessment.customer_id == customer_id,
FinRiskAssessment.valid_until >= date.today(),
)
.order_by(
FinRiskAssessment.assessment_date.desc(),
FinRiskAssessment.id.desc(),
)
.limit(1)
)
class CustomerProfileRepo(BaseRepository):
"""客户画像仓储。注意:主键是 customer_id 而非 id,不适用基类 get(pk)。"""
+4
View File
@@ -22,3 +22,7 @@ class SensitiveWordRepo(BaseRepository):
)
).all()
)
async def list_active_words(self) -> list[str]:
"""Return enabled words as strings for Agent content scanning."""
return [item if isinstance(item, str) else item.word for item in await self.list_active()]
+9
View File
@@ -10,6 +10,15 @@ from repositories.base import BaseRepository
class SysMessageRepo(BaseRepository):
model = SysMessage
async def get_by_biz_id(self, biz_id: str, *, user_id: int) -> SysMessage | None:
"""按收件人和业务号查询已写入的站内信,供发送重试幂等使用。"""
return await self.db.scalar(
select(SysMessage).where(
SysMessage.biz_id == biz_id,
SysMessage.user_id == user_id,
)
)
async def list_by_user(
self, user_id: int, limit: int = 100, offset: int = 0
) -> list[SysMessage]: