|
1 | | -"""聊天管理器""" |
| 1 | +"""聊天管理器 |
2 | 2 |
|
| 3 | +Business logic for chat operations using SQLModel and FastAPI patterns. |
| 4 | +""" |
| 5 | + |
| 6 | +__all__ = ["ChatManager"] |
| 7 | + |
| 8 | +import json |
3 | 9 | import typing |
4 | | -from typing import Annotated as Anno, Literal as Lit, Optional as Opt |
5 | | - |
6 | | -from dal import DefaultRedis |
7 | | -from blue_firmament import Method, listen_to |
8 | | -from blue_firmament.exceptions import Forbidden, NotFound, ParamsInvalid, Conflict |
9 | | -from blue_firmament.log import get_logger |
10 | | -from blue_firmament.manager import CommonManager, PresetHandlerConfig |
11 | | -from blue_firmament.scheme.converter import IntConverter |
12 | | -from blue_firmament.task import TaskStatus |
13 | | -from blue_firmament.task.result import StreamingBody, PlainTextBody |
| 10 | +from typing import Optional as Opt |
| 11 | + |
| 12 | +import sqlmodel |
| 13 | + |
| 14 | +from core.engine import SessionLocal |
14 | 15 | from account.schemas import AccountRef |
15 | | -from ..schemas.chat import Chat, ChatRef, ChatType |
| 16 | +from ..schemas.chat import Chat, ChatRef, ChatType, ChatStatus |
16 | 17 | from ..schemas.message import Message |
17 | | -from main.schemas.partner_request import PartnerRequest |
18 | 18 |
|
19 | 19 | if typing.TYPE_CHECKING: |
| 20 | + from main.schemas.partner_request import PartnerRequest |
20 | 21 | from main.schemas.partner_request.application import PartnerApplication |
21 | 22 |
|
22 | | -LOGGER = get_logger(__name__) |
23 | 23 |
|
| 24 | +class ChatManager: |
| 25 | + """Chat business logic manager. |
| 26 | +
|
| 27 | + Provides methods for chat operations without BlueFirmament dependencies. |
| 28 | + Uses SQLModel sessions directly. |
| 29 | + """ |
| 30 | + |
| 31 | + @classmethod |
| 32 | + def get(cls, chat_id: ChatRef) -> Opt[Chat]: |
| 33 | + """Get a chat by ID.""" |
| 34 | + with SessionLocal() as db: |
| 35 | + return db.get(Chat, chat_id) |
24 | 36 |
|
25 | | -class ChatManager( |
26 | | - CommonManager[Chat, ChatRef], |
27 | | - scheme_cls=Chat, |
28 | | - manager_name="chat", |
29 | | - path_prefix="chat", |
30 | | - preset_handler_config=PresetHandlerConfig(get=True), |
31 | | -): |
32 | | - async def _must_be_member( |
33 | | - self, chat_id: Opt[ChatRef] = None, account_id: Opt[AccountRef] = None |
34 | | - ) -> None: |
35 | | - """必须为聊天成员 |
| 37 | + @classmethod |
| 38 | + def get_members(cls, chat_id: ChatRef) -> set[AccountRef]: |
| 39 | + """获取聊天成员列表 |
36 | 40 |
|
37 | 41 | :param chat_id: 聊天 ID |
38 | | - :raise Forbidden: 如果不是成员 |
| 42 | + :return: 成员 ID 集合 |
39 | 43 | """ |
40 | | - account_id = account_id or AccountRef(self._operator.id) |
41 | | - chat = await self._get_scheme(_id=chat_id) |
42 | | - if account_id not in (await chat.get_members()): |
43 | | - raise Forbidden("must be a member of the chat") |
44 | | - |
45 | | - @listen_to(Method.GET, "/{chat_id}/messages") |
46 | | - async def get_history( |
47 | | - self, |
| 44 | + with SessionLocal() as db: |
| 45 | + chat = db.get(Chat, chat_id) |
| 46 | + if not chat: |
| 47 | + return set() |
| 48 | + if chat.members is None: |
| 49 | + return set() |
| 50 | + members_list = json.loads(chat.members) |
| 51 | + return set(members_list) |
| 52 | + |
| 53 | + @classmethod |
| 54 | + def is_member(cls, chat_id: ChatRef, account_id: AccountRef) -> bool: |
| 55 | + """检查用户是否为聊天成员 |
| 56 | +
|
| 57 | + :param chat_id: 聊天 ID |
| 58 | + :param account_id: 账号 ID |
| 59 | + :return: 是否为成员 |
| 60 | + """ |
| 61 | + members = cls.get_members(chat_id) |
| 62 | + return account_id in members |
| 63 | + |
| 64 | + @classmethod |
| 65 | + def get_chat_messages( |
| 66 | + cls, |
48 | 67 | chat_id: ChatRef, |
49 | | - start: Anno[int, IntConverter(ge=0)] = 0, |
50 | | - offset: Anno[int, IntConverter(ge=1, le=12)] = 6, |
| 68 | + start: int = 0, |
| 69 | + offset: int = 6, |
51 | 70 | desc: bool = True, |
52 | | - ) -> typing.Tuple[Message, ...]: |
| 71 | + ) -> list[Message]: |
53 | 72 | """获取聊天历史消息 |
54 | 73 |
|
| 74 | + :param chat_id: 聊天 ID |
| 75 | + :param start: 起始位置 |
| 76 | + :param offset: 获取数量 |
55 | 77 | :param desc: 是否降序排列 |
56 | | -
|
57 | | - 条件: |
58 | | - - 必须是聊天的成员 |
| 78 | + :return: 消息列表 |
59 | 79 | """ |
60 | | - await self._must_be_member(chat_id=chat_id) |
61 | | - from .message import BaseMessageManager |
62 | | - |
63 | | - return await BaseMessageManager(self).get_chat_messages( |
64 | | - chat_id=chat_id, start=start, offset=offset, desc=desc |
65 | | - ) |
66 | | - |
67 | | - async def get_members(self, chat_id: Opt[ChatRef] = None) -> set[AccountRef]: |
68 | | - """获取聊天成员列表""" |
69 | | - self._scheme = await self._get_scheme(chat_id) |
70 | | - return await self._scheme.get_members() |
71 | | - |
72 | | - @listen_to(Method.GET, "/mine") |
73 | | - async def get_mine( |
74 | | - self, chat_type: Opt[ChatType] = None, return_in: Lit["id", "full"] = "id" |
75 | | - ) -> tuple[ChatRef | Chat, ...]: |
76 | | - """获取我的聊天 |
77 | | -
|
78 | | - :param chat_type: 聊天类型 |
79 | | - :param return_in: 返回类型,默认为 ID 列表 |
80 | | - - "id": 仅返回聊天 ID 列表 |
81 | | - - "full": 返回完整的聊天对象列表 |
82 | | - :return: 聊天列表 |
83 | | -
|
84 | | - Docs |
85 | | - ---- |
86 | | - - `APIFOX <https://app.apifox.com/link/project/4406548/apis/api-275041592>`_ |
| 80 | + with SessionLocal() as db: |
| 81 | + statement = sqlmodel.select(Message).where(Message.chat == chat_id) |
| 82 | + if desc: |
| 83 | + statement = statement.order_by(Message.id.desc()) |
| 84 | + else: |
| 85 | + statement = statement.order_by(Message.id) |
| 86 | + statement = statement.offset(start).limit(offset) |
| 87 | + messages = db.exec(statement).all() |
| 88 | + return list(messages) |
| 89 | + |
| 90 | + @classmethod |
| 91 | + def get_user_chats( |
| 92 | + cls, |
| 93 | + user_id: AccountRef, |
| 94 | + chat_type: Opt[ChatType] = None, |
| 95 | + ) -> list[ChatRef]: |
| 96 | + """获取用户的聊天列表 |
| 97 | +
|
| 98 | + :param user_id: 用户 ID |
| 99 | + :param chat_type: 聊天类型过滤 |
| 100 | + :return: 聊天 ID 列表 |
87 | 101 | """ |
88 | | - query_coms = (Chat.type.equals(chat_type),) if chat_type else () |
89 | | - if return_in == "full": |
90 | | - return await self._dao.select(*query_coms) |
91 | | - elif return_in == "id": |
92 | | - return await self._dao.select_field(Chat._id, *query_coms) |
93 | | - raise ParamsInvalid("unsupported", return_in=return_in) |
94 | | - |
95 | | - @listen_to(Method.PUT, "/direct_message/{to_id}") |
96 | | - async def create_dm_chat( |
97 | | - self, |
| 102 | + with SessionLocal() as db: |
| 103 | + statement = sqlmodel.select(Chat.id).where(Chat.created_by == user_id) |
| 104 | + if chat_type: |
| 105 | + statement = statement.where(Chat.type == chat_type.value) |
| 106 | + results = db.exec(statement).all() |
| 107 | + return list(results) |
| 108 | + |
| 109 | + @classmethod |
| 110 | + def create_dm_chat( |
| 111 | + cls, |
| 112 | + from_id: AccountRef, |
98 | 113 | to_id: AccountRef, |
99 | | - from_id: Opt[AccountRef] = None, |
100 | | - ) -> Chat: |
| 114 | + ) -> tuple[Chat, bool]: |
101 | 115 | """创建私信聊天 |
102 | 116 |
|
103 | 117 | 如果两者已经存在私信聊天,则返回已存在的聊天 |
104 | 118 |
|
105 | | - :param to_id: 私聊对象 |
106 | | - :param from_id: 私聊发起者 |
| 119 | + :param from_id: 发起者 ID |
| 120 | + :param to_id: 目标用户 ID |
| 121 | + :return: (聊天对象, 是否新创建) |
| 122 | + :raises ValueError: 如果试图与自己创建私信 |
107 | 123 | """ |
108 | | - from_id = from_id or AccountRef(self._operator.id) # TODO use param getter (resolver) |
109 | | - |
110 | 124 | if from_id == to_id: |
111 | | - raise Conflict("cannot create a direct message chat with yourself") |
| 125 | + raise ValueError("cannot create a direct message chat with yourself") |
112 | 126 |
|
113 | | - try: |
114 | | - self._scheme = await self._dao.select_one( |
115 | | - Chat.type.equals(ChatType.DIRECT_MESSAGE), |
116 | | - Chat.members.contains(to_id, from_id), |
| 127 | + with SessionLocal() as db: |
| 128 | + # 尝试查找已存在的私信聊天 |
| 129 | + statement = sqlmodel.select(Chat).where( |
| 130 | + Chat.type == ChatType.DIRECT_MESSAGE.value |
117 | 131 | ) |
118 | | - except NotFound: |
119 | | - # 创建私信聊天 |
120 | | - self._scheme = await self.insert( |
121 | | - Chat( |
122 | | - _task_context=self, |
123 | | - _id=ChatRef(0), |
124 | | - type=ChatType.DIRECT_MESSAGE, |
125 | | - created_by=from_id, |
126 | | - members={to_id, from_id}, |
127 | | - ) |
| 132 | + chats = db.exec(statement).all() |
| 133 | + |
| 134 | + for chat in chats: |
| 135 | + if chat.members: |
| 136 | + members = set(json.loads(chat.members)) |
| 137 | + if {from_id, to_id} == members: |
| 138 | + return chat, False |
| 139 | + |
| 140 | + # 创建新的私信聊天 |
| 141 | + members_json = json.dumps([from_id, to_id]) |
| 142 | + chat = Chat( |
| 143 | + type=ChatType.DIRECT_MESSAGE.value, |
| 144 | + created_by=from_id, |
| 145 | + members=members_json, |
128 | 146 | ) |
129 | | - self._task_result.status = TaskStatus.CREATED |
| 147 | + db.add(chat) |
| 148 | + db.commit() |
| 149 | + db.refresh(chat) |
| 150 | + return chat, True |
130 | 151 |
|
131 | | - return self._scheme |
132 | | - |
133 | | - async def create_pr_chat(self, partner_request: PartnerRequest) -> Chat: |
| 152 | + @classmethod |
| 153 | + def create_pr_chat(cls, partner_request: "PartnerRequest") -> Chat: |
134 | 154 | """创建搭子请求群聊 |
135 | 155 |
|
136 | | - - 创建者是搭子请求的创建者 |
137 | | - - 成员列表为 None |
| 156 | + :param partner_request: 搭子请求 |
| 157 | + :return: 创建的聊天 |
138 | 158 | """ |
139 | | - return await self.insert( |
140 | | - Chat( |
141 | | - _id=ChatRef(0), |
142 | | - type=ChatType.PARTNER_REQUEST, |
| 159 | + with SessionLocal() as db: |
| 160 | + chat = Chat( |
| 161 | + type=ChatType.PARTNER_REQUEST.value, |
143 | 162 | created_by=partner_request.created_by, |
144 | 163 | members=None, |
145 | 164 | ) |
146 | | - ) |
147 | | - |
148 | | - async def create_partner_application_chat( |
149 | | - self, |
| 165 | + db.add(chat) |
| 166 | + db.commit() |
| 167 | + db.refresh(chat) |
| 168 | + return chat |
| 169 | + |
| 170 | + @classmethod |
| 171 | + def create_partner_application_chat( |
| 172 | + cls, |
150 | 173 | application: "PartnerApplication", |
151 | 174 | pr_chat_id: ChatRef, |
152 | 175 | ) -> Chat: |
153 | 176 | """创建搭子申请群聊 |
154 | 177 |
|
155 | | - 1. 创建搭子申请群聊 |
156 | | - 3. 链接搭子申请群聊到搭子请求群聊的子群聊中 |
157 | | -
|
158 | 178 | :param application: 搭子申请 |
159 | 179 | :param pr_chat_id: 搭子请求群聊 ID |
| 180 | + :return: 创建的聊天 |
160 | 181 | """ |
161 | | - # 1. create chat |
162 | | - chat = await self.insert( |
163 | | - Chat( |
164 | | - _id=ChatRef(0), |
165 | | - type=ChatType.PARTNER_APPLICATION, |
| 182 | + with SessionLocal() as db: |
| 183 | + chat = Chat( |
| 184 | + type=ChatType.PARTNER_APPLICATION.value, |
166 | 185 | created_by=application.applicant, |
167 | | - parent=pr_chat_id, # 链接到搭子请求群聊 |
| 186 | + parent=pr_chat_id, |
168 | 187 | members=None, |
169 | 188 | ) |
170 | | - ) |
171 | | - |
172 | | - return chat |
173 | | - |
174 | | - @listen_to(Method.GET, "/unread") |
175 | | - async def get_my_unread( |
176 | | - self, |
177 | | - ) -> StreamingBody: |
178 | | - """持续获取未读消息""" |
179 | | - redis = DefaultRedis() |
180 | | - listen_to_queue = f"unread_messages:{self._operator.id}" |
181 | | - |
182 | | - # await redis.subscribe(*listen_to_channels) |
183 | | - |
184 | | - stop = False |
185 | | - |
186 | | - async def _get_my_unread() -> StreamingBody.GeneratorType: |
187 | | - while not stop: |
188 | | - # message = await redis.get_message() |
189 | | - message_id = await redis.pop(listen_to_queue) |
190 | | - yield PlainTextBody(message_id.decode("utf-8")) |
191 | | - self._logger.info("Stop listening for unread messages") |
192 | | - |
193 | | - async def _cleanup() -> None: |
194 | | - nonlocal stop |
195 | | - stop = True |
196 | | - await redis.close() |
197 | | - |
198 | | - return StreamingBody( |
199 | | - generator=_get_my_unread(), |
200 | | - cleanup=_cleanup, |
201 | | - ) |
202 | | - |
203 | | - async def close(self, chat_id: Opt[ChatRef] = None) -> Chat: |
204 | | - """关闭群聊""" |
205 | | - self._scheme = await self._get_scheme(_id=chat_id) |
206 | | - self._scheme.status = self._scheme.status.to_closed() |
207 | | - return await self._update_scheme() |
| 189 | + db.add(chat) |
| 190 | + db.commit() |
| 191 | + db.refresh(chat) |
| 192 | + return chat |
| 193 | + |
| 194 | + @classmethod |
| 195 | + def close_chat(cls, chat_id: ChatRef) -> Opt[Chat]: |
| 196 | + """关闭群聊 |
| 197 | +
|
| 198 | + :param chat_id: 聊天 ID |
| 199 | + :return: 更新后的聊天对象,如果不存在则返回 None |
| 200 | + """ |
| 201 | + with SessionLocal() as db: |
| 202 | + chat = db.get(Chat, chat_id) |
| 203 | + if not chat: |
| 204 | + return None |
| 205 | + chat.status = ChatStatus.CLOSED.value |
| 206 | + db.add(chat) |
| 207 | + db.commit() |
| 208 | + db.refresh(chat) |
| 209 | + return chat |
| 210 | + |
0 commit comments