feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
*,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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())
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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)。"""
|
||||
|
||||
@@ -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()]
|
||||
|
||||
@@ -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]:
|
||||
|
||||
Reference in New Issue
Block a user