diff --git a/app/api/routers/rbac_investor.py b/app/api/routers/rbac_investor.py index e3c0358..3a079d2 100644 --- a/app/api/routers/rbac_investor.py +++ b/app/api/routers/rbac_investor.py @@ -192,18 +192,11 @@ async def create_roadshow( db: Database = Depends(get_db), actor: dict = Depends(require_roles(*_ROADSHOW_ROLES)), ): - if not req.title.strip(): - raise HTTPException(status_code=400, detail="标题不能为空") - if req.end_at and req.start_at and req.end_at < req.start_at: - raise HTTPException(status_code=400, detail="结束时间不能早于开始时间") + from ...services.roadshow_service import RoadshowService - need_review = _need_review(actor, req, db) - fields = {**req.model_dump(exclude_none=True), "publisher_id": actor["id"], - "publisher_role": actor.get("role", "investor"), - "status": _pub_status(need_review), "need_review": need_review} - rs = await db.roadshows.create(fields) + rs = await RoadshowService(db).create(req, actor) await write_audit(db, action="roadshow.publish", resource="roadshow", resource_id=rs["id"], - detail=f"{rs['title']} need_review={need_review}", user=actor, request=request) + detail=f"{rs['title']} need_review={rs.get('need_review')}", user=actor, request=request) return rs diff --git a/app/api/routers/rbac_org.py b/app/api/routers/rbac_org.py index 79d189d..aef05c7 100644 --- a/app/api/routers/rbac_org.py +++ b/app/api/routers/rbac_org.py @@ -107,13 +107,6 @@ async def my_orgs( user: dict = Depends(require_roles("enterprise", "carrier", "provider", "investor")), ): """返回当前账号挂靠的所有机构及成员角色(支持一账号多机构)。""" - out = [] - for org in await db.orgs.all(): - member = None - for m in await db.org_members.list_for_org(org["id"]): - if m["user_id"] == user["id"]: - member = m - break - if member: - out.append({"org": org, "member_role": member["role"], "is_admin": member["is_admin"]}) - return {"items": out} + from ...services.org_service import OrgService + + return await OrgService(db).my_orgs(user["id"]) diff --git a/app/infrastructure/repositories.py b/app/infrastructure/repositories.py index 6c006df..1845d19 100644 --- a/app/infrastructure/repositories.py +++ b/app/infrastructure/repositories.py @@ -527,6 +527,13 @@ class RegionRepository: return {"id": r.id, "name": r.name, "level": r.level, "parent_id": r.parent_id} if r else None async def all(self) -> list[dict]: + return [_to_dict_org(o) for o in await self.session.scalars(select(Organization).order_by(Organization.created_at))] + + async def list_by_ids(self, org_ids: list[str]) -> list[dict]: + if not org_ids: + return [] + rows = await self.session.scalars(select(Organization).where(Organization.id.in_(org_ids))) + return [_to_dict_org(o) for o in rows] return [ {"id": r.id, "name": r.name, "level": r.level, "parent_id": r.parent_id} for r in await self.session.scalars(select(Region)) @@ -585,6 +592,13 @@ class OrgRepository: return self._to_dict(o) if o else None async def all(self) -> list[dict]: + return [_to_dict_org(o) for o in await self.session.scalars(select(Organization).order_by(Organization.created_at))] + + async def list_by_ids(self, org_ids: list[str]) -> list[dict]: + if not org_ids: + return [] + rows = await self.session.scalars(select(Organization).where(Organization.id.in_(org_ids))) + return [_to_dict_org(o) for o in rows] return [self._to_dict(o) for o in await self.session.scalars(select(Organization))] async def list_by_region(self, region_ids: list[str]) -> list[dict]: @@ -899,6 +913,13 @@ class ConfigRepository: self.session = session async def all(self) -> list[dict]: + return [_to_dict_org(o) for o in await self.session.scalars(select(Organization).order_by(Organization.created_at))] + + async def list_by_ids(self, org_ids: list[str]) -> list[dict]: + if not org_ids: + return [] + rows = await self.session.scalars(select(Organization).where(Organization.id.in_(org_ids))) + return [_to_dict_org(o) for o in rows] return [ {"key": c.key, "value": c.value, "description": c.description, "updated_at": c.updated_at} @@ -1678,6 +1699,14 @@ class OrganizationMemberRepository: def __init__(self, session: AsyncSession): self.session = session + async def list_for_user(self, user_id: str) -> list[dict]: + rows = await self.session.scalars(select(OrganizationMember).where(OrganizationMember.user_id == user_id)) + return [self._to_dict(m) for m in rows] + + async def get(self, org_id: str, user_id: str) -> dict | None: + m = await self.session.get(OrganizationMember, (org_id, user_id)) + return self._to_dict(m) if m else None + async def list_for_org(self, org_id: str) -> list[dict]: rows = (await self.session.scalars( select(OrganizationMember).where(OrganizationMember.org_id == org_id) diff --git a/app/services/org_service.py b/app/services/org_service.py new file mode 100644 index 0000000..e590a01 --- /dev/null +++ b/app/services/org_service.py @@ -0,0 +1,40 @@ +# -*- coding: utf-8 -*- +"""业务层 · 组织机构服务(成员管理,消除 N+1)。""" +from __future__ import annotations + +from fastapi import HTTPException + +from ..infrastructure.repositories import Database + + +class OrgService: + """机构/组织成员管理。""" + + def __init__(self, db: Database): + self.db = db + + async def my_orgs(self, user_id: str) -> dict: + """返回当前账号挂靠的所有机构及成员角色(支持一账号多机构,批量查避免 N+1)。""" + members = await self.db.org_members.list_for_user(user_id) + org_ids = [m["org_id"] for m in members] + orgs = {o["id"]: o for o in await self.db.orgs.list_by_ids(org_ids)} if org_ids else {} + return {"items": [ + {"org": orgs[m["org_id"]], "member_role": m["role"], "is_admin": m["is_admin"]} + for m in members if m["org_id"] in orgs + ]} + + async def add_member(self, org_id: str, username: str, role: str, actor: dict) -> dict: + """添加机构成员(须机构 admin,用户名解析为 user_id)。""" + org = await self.db.orgs.get(org_id) + if org is None: + raise HTTPException(status_code=404, detail="机构不存在") + actor_members = await self.db.org_members.list_for_org(org_id) + if not any(m["user_id"] == actor["id"] and m["is_admin"] for m in actor_members): + raise HTTPException(status_code=403, detail="仅机构管理员可添加成员") + user = await self.db.users.get_by_username(username) + if user is None: + raise HTTPException(status_code=404, detail="用户不存在") + existing = await self.db.org_members.get(org_id, user["id"]) + if existing: + return existing + return await self.db.org_members.add_member(org_id, user["id"], role) diff --git a/app/services/roadshow_service.py b/app/services/roadshow_service.py new file mode 100644 index 0000000..40b7d4f --- /dev/null +++ b/app/services/roadshow_service.py @@ -0,0 +1,42 @@ +# -*- coding: utf-8 -*- +"""业务层 · 路演服务(发布审核规则)。""" +from __future__ import annotations + +from fastapi import HTTPException + +from ..infrastructure.repositories import Database + + +class RoadshowService: + """路演发布:按发布主体与范围判定是否需审核。""" + + def __init__(self, db: Database): + self.db = db + + def _need_review(self, user: dict, req) -> bool: + role = user.get("role") + if role == "operator": + return False + if role == "government": + scope = user.get("scope_region_ids", []) + if req.scope_type == "region" and req.region_id and req.region_id in scope: + return False # 自身权限范围内,免审 + return True # 超出范围,需审核 + # 投资机构/载体/企业:一律需审核 + return True + + async def create(self, req, actor: dict) -> dict: + """创建路演(含审核判定)。""" + if not req.title.strip(): + raise HTTPException(status_code=400, detail="标题不能为空") + if req.end_at and req.start_at and req.end_at < req.start_at: + raise HTTPException(status_code=400, detail="结束时间不能早于开始时间") + need_review = self._need_review(actor, req) + fields = { + **req.model_dump(exclude_none=True), + "publisher_id": actor["id"], + "publisher_role": actor.get("role", "investor"), + "status": "submitted" if need_review else "published", + "need_review": need_review, + } + return await self.db.roadshows.create(fields)