fix(message): make inbox read-status writes idempotent under concurrency

batch_mark_as_read / mark_as_read used check-then-insert with no protection
against a concurrent transaction inserting the same (message_id, user_id) in
the SELECT->INSERT window, tripping ix_inbox_message_read_msg_user. batch used
one add_all + one commit, so a single conflict rolled back the whole batch —
every message stayed unread and the API 500'd (prod trace gkpj / DM -6602).

- batch_mark_as_read: on IntegrityError, rollback, re-read what now exists and
  insert only the still-missing ids (dialect-agnostic, same pattern as
  ChannelInfoSourceRepositoryImpl.batch_add); dedupe input ids too.
- mark_as_read: on conflict, rollback and return the row that won the race;
  re-raise only when it is not a duplicate-key race.
- handle_approval: guard the post-action auto-read (narrow IntegrityError
  catch) so a read-record race never surfaces the finished approval as a 500.

Adds deterministic concurrency regression tests that inject a competing insert
into the exact race window.
This commit is contained in:
GuoQing Zhang
2026-07-06 16:10:52 +08:00
parent e5543750a8
commit ca72ae8bb1
3 changed files with 339 additions and 141 deletions
@@ -1,8 +1,7 @@
import logging
from typing import List
from sqlalchemy import func
from sqlmodel import select, col
from sqlalchemy.exc import IntegrityError
from sqlmodel import col, select
from sqlmodel.ext.asyncio.session import AsyncSession
from bisheng.common.repositories.implementations.base_repository_impl import BaseRepositoryImpl
@@ -25,42 +24,74 @@ class InboxMessageReadRepositoryImpl(BaseRepositoryImpl[InboxMessageRead, int],
return existing
record = InboxMessageRead(message_id=message_id, user_id=user_id)
return await self.save(record)
try:
return await self.save(record)
except IntegrityError:
# A concurrent request inserted the same (message_id, user_id) between
# our find_one and save, tripping the ix_inbox_message_read_msg_user
# unique index. Roll back and return the row that won the race so
# marking-as-read stays idempotent instead of surfacing a 500.
await self.session.rollback()
existing = await self.find_one(message_id=message_id, user_id=user_id)
if existing:
return existing
raise # not a duplicate-key race — a real failure, surface it
async def batch_mark_as_read(self, message_ids: List[int], user_id: int) -> int:
"""Batch mark messages as read. Returns number of newly marked records."""
async def batch_mark_as_read(self, message_ids: list[int], user_id: int) -> int:
"""Batch mark messages as read. Returns number of newly marked records.
Idempotent under concurrency: if a competing request inserts the same
(message_id, user_id) in the SELECT->INSERT window, one unique-constraint
conflict would otherwise roll back the *whole* batch (leaving every
message unread and 500ing the API). Instead we roll back, re-read what
now exists and insert only the still-missing ids. Same dialect-agnostic
pattern as ChannelInfoSourceRepositoryImpl.batch_add.
"""
if not message_ids:
return 0
# Find which messages are already read
new_records = await self._build_unread_records(message_ids, user_id)
if not new_records:
return 0
try:
self.session.add_all(new_records)
await self.session.commit()
return len(new_records)
except IntegrityError:
await self.session.rollback()
remaining = await self._build_unread_records(message_ids, user_id)
if remaining:
self.session.add_all(remaining)
await self.session.commit()
return len(remaining)
async def _build_unread_records(self, message_ids: list[int], user_id: int) -> list[InboxMessageRead]:
"""Build read-records for ids not yet read by the user.
Dedupes the input (``dict.fromkeys`` preserves order) so duplicate ids in
a single request cannot self-collide on the unique index.
"""
query = select(InboxMessageRead.message_id).where(
InboxMessageRead.user_id == user_id,
col(InboxMessageRead.message_id).in_(message_ids)
col(InboxMessageRead.message_id).in_(message_ids),
)
result = await self.session.exec(query)
already_read_ids = set(result.all())
# Create records for unread messages only
new_records = []
for msg_id in message_ids:
if msg_id not in already_read_ids:
new_records.append(InboxMessageRead(message_id=msg_id, user_id=user_id))
if new_records:
self.session.add_all(new_records)
await self.session.commit()
return len(new_records)
return [
InboxMessageRead(message_id=msg_id, user_id=user_id)
for msg_id in dict.fromkeys(message_ids)
if msg_id not in already_read_ids
]
async def is_read(self, message_id: int, user_id: int) -> bool:
"""Check if a message has been read by a specific user."""
existing = await self.find_one(message_id=message_id, user_id=user_id)
return existing is not None
async def get_read_message_ids(self, user_id: int) -> List[int]:
async def get_read_message_ids(self, user_id: int) -> list[int]:
"""Get all read message IDs for a specific user."""
query = select(InboxMessageRead.message_id).where(
InboxMessageRead.user_id == user_id
)
query = select(InboxMessageRead.message_id).where(InboxMessageRead.user_id == user_id)
result = await self.session.exec(query)
return list(result.all())
@@ -1,28 +1,31 @@
import copy
import logging
from typing import List, Optional, Any, Dict
from typing import Any
from sqlalchemy.exc import IntegrityError
from bisheng.common.dependencies.user_deps import UserPayload
from bisheng.common.errcode.message import (
MessageAlreadyApprovedError,
MessageNotFoundError,
MessagePermissionDeniedError,
MessageAlreadyApprovedError,
)
from bisheng.message.domain.models.inbox_message import InboxMessage, MessageTypeEnum, MessageStatusEnum
from bisheng.database.models.user_group import UserGroupDao
from bisheng.message.domain.models.inbox_message import InboxMessage, MessageStatusEnum, MessageTypeEnum
from bisheng.message.domain.repositories.interfaces.inbox_message_read_repository import InboxMessageReadRepository
from bisheng.message.domain.repositories.interfaces.inbox_message_repository import InboxMessageRepository
from bisheng.message.domain.schemas.message_schema import (
ApprovalActionEnum,
MessageContentItem,
MessageItemResponse,
MessagePageResponse,
TabTypeEnum,
UnreadCountResponse,
ApprovalActionEnum,
TabTypeEnum, MessageContentItem,
)
from bisheng.message.domain.services.approval_handler import ApprovalHandler
from bisheng.message.domain.services.notification_content import infer_action_code
from bisheng.database.models.user_group import UserGroupDao
from bisheng.user.domain.models.user import UserDao
from bisheng.notification.forwarder import maybe_forward_external
from bisheng.user.domain.models.user import UserDao
logger = logging.getLogger(__name__)
@@ -44,14 +47,14 @@ class MessageService:
"""Service layer for in-app messaging (inbox) operations."""
def __init__(
self,
message_repository: 'InboxMessageRepository',
message_read_repository: 'InboxMessageReadRepository',
approval_handlers: Optional[List[ApprovalHandler]] = None,
self,
message_repository: "InboxMessageRepository",
message_read_repository: "InboxMessageReadRepository",
approval_handlers: list[ApprovalHandler] | None = None,
):
self.message_repository = message_repository
self.message_read_repository = message_read_repository
self._handler_map: Dict[str, ApprovalHandler] = {}
self._handler_map: dict[str, ApprovalHandler] = {}
for handler in approval_handlers or []:
action_code = handler.get_action_code()
@@ -60,13 +63,13 @@ class MessageService:
self._handler_map[action_code] = handler
async def send_message(
self,
content: List[Dict[str, Any]],
sender: int,
message_type: MessageTypeEnum,
receiver: List[int],
status: MessageStatusEnum = MessageStatusEnum.WAIT_APPROVE,
action_code: Optional[str] = None,
self,
content: list[dict[str, Any]],
sender: int,
message_type: MessageTypeEnum,
receiver: list[int],
status: MessageStatusEnum = MessageStatusEnum.WAIT_APPROVE,
action_code: str | None = None,
) -> InboxMessage:
"""Create and save a new inbox message.
@@ -84,7 +87,10 @@ class MessageService:
saved_message = await self.message_repository.save(message)
logger.info(
"Inbox message sent: id=%s, type=%s, sender=%s, receivers=%s",
saved_message.id, message_type.value, sender, receiver,
saved_message.id,
message_type.value,
sender,
receiver,
)
# E+ forwarding hook — sync, lightweight; HTTP fires via asyncio task
@@ -97,13 +103,13 @@ class MessageService:
return saved_message
async def get_message_list(
self,
login_user: UserPayload,
tab: TabTypeEnum = TabTypeEnum.ALL,
only_unread: bool = False,
keyword: Optional[str] = None,
page: int = 1,
page_size: int = 20,
self,
login_user: UserPayload,
tab: TabTypeEnum = TabTypeEnum.ALL,
only_unread: bool = False,
keyword: str | None = None,
page: int = 1,
page_size: int = 20,
) -> MessagePageResponse:
"""Get paginated message list for the current user with read status annotation."""
# 1. Get read message IDs for this user
@@ -141,7 +147,7 @@ class MessageService:
# 5. Batch query sender user names
sender_ids = list({m.sender for m in messages})
sender_map: Dict[int, str] = {}
sender_map: dict[int, str] = {}
if sender_ids:
users = await UserDao.aget_user_by_ids(sender_ids)
sender_map = {u.user_id: u.user_name for u in users}
@@ -153,76 +159,75 @@ class MessageService:
# 7. Build response
items = []
for msg in messages:
items.append(MessageItemResponse(
id=msg.id,
content=self._enrich_message_content_with_group_names(msg.content, user_group_name_map),
sender=msg.sender,
sender_name=sender_map.get(msg.sender),
message_type=msg.message_type.value,
status=msg.status.value,
action_code=msg.action_code,
operator_user_id=msg.operator_user_id,
is_read=msg.id in read_set,
create_time=msg.create_time,
update_time=msg.update_time,
))
items.append(
MessageItemResponse(
id=msg.id,
content=self._enrich_message_content_with_group_names(msg.content, user_group_name_map),
sender=msg.sender,
sender_name=sender_map.get(msg.sender),
message_type=msg.message_type.value,
status=msg.status.value,
action_code=msg.action_code,
operator_user_id=msg.operator_user_id,
is_read=msg.id in read_set,
create_time=msg.create_time,
update_time=msg.update_time,
)
)
return MessagePageResponse(data=items, total=total)
@staticmethod
def _extract_content_user_ids(messages: List[InboxMessage]) -> List[int]:
def _extract_content_user_ids(messages: list[InboxMessage]) -> list[int]:
"""Extract distinct user IDs from content items with type=user."""
user_ids: set[int] = set()
for message in messages:
for item in message.content or []:
if item.get('type') != 'user':
if item.get("type") != "user":
continue
metadata = item.get('metadata') or {}
user_id = metadata.get('user_id')
metadata = item.get("metadata") or {}
user_id = metadata.get("user_id")
if isinstance(user_id, int):
user_ids.add(user_id)
return list(user_ids)
@staticmethod
async def _build_user_group_name_map(user_ids: List[int]) -> Dict[int, List[str]]:
async def _build_user_group_name_map(user_ids: list[int]) -> dict[int, list[str]]:
"""Build a map from user_id to the user's group names."""
if not user_ids:
return {}
user_groups_map = await UserGroupDao.aget_user_groups_batch(user_ids)
return {
user_id: [group.group_name for group in groups]
for user_id, groups in user_groups_map.items()
}
return {user_id: [group.group_name for group in groups] for user_id, groups in user_groups_map.items()}
@staticmethod
def _enrich_message_content_with_group_names(
content: List[Dict[str, Any]],
user_group_name_map: Dict[int, List[str]],
) -> List[Dict[str, Any]]:
content: list[dict[str, Any]],
user_group_name_map: dict[int, list[str]],
) -> list[dict[str, Any]]:
"""Attach group_names into metadata for content items with type=user."""
enriched_content = []
for item in content or []:
if item.get('type') != 'user':
if item.get("type") != "user":
enriched_content.append(item)
continue
metadata = item.get('metadata')
metadata = item.get("metadata")
if not isinstance(metadata, dict):
enriched_content.append(item)
continue
user_id = metadata.get('user_id')
user_id = metadata.get("user_id")
if not isinstance(user_id, int):
enriched_content.append(item)
continue
new_item = dict(item)
new_metadata = dict(metadata)
new_metadata['group_names'] = user_group_name_map.get(user_id, [])
new_item['metadata'] = new_metadata
new_metadata["group_names"] = user_group_name_map.get(user_id, [])
new_item["metadata"] = new_metadata
enriched_content.append(new_item)
return enriched_content
@@ -250,7 +255,7 @@ class MessageService:
return UnreadCountResponse(total=total, notify=notify_count, approve=approve_count)
async def mark_as_read(self, message_ids: List[int], login_user: UserPayload) -> int:
async def mark_as_read(self, message_ids: list[int], login_user: UserPayload) -> int:
"""Mark specific messages as read for the current user."""
return await self.message_read_repository.batch_mark_as_read(message_ids, login_user.user_id)
@@ -270,10 +275,10 @@ class MessageService:
return await self.message_read_repository.batch_mark_as_read(unread_ids, login_user.user_id)
async def handle_approval(
self,
message_id: int,
action: ApprovalActionEnum,
login_user: UserPayload,
self,
message_id: int,
action: ApprovalActionEnum,
login_user: UserPayload,
) -> InboxMessage:
"""
Handle approval action (agree/reject) on an approval message.
@@ -294,11 +299,7 @@ class MessageService:
raise MessageAlreadyApprovedError()
# 4. Determine new status
new_status = (
MessageStatusEnum.APPROVED
if action == ApprovalActionEnum.AGREE
else MessageStatusEnum.REJECTED
)
new_status = MessageStatusEnum.APPROVED if action == ApprovalActionEnum.AGREE else MessageStatusEnum.REJECTED
original_content = copy.deepcopy(message.content)
action_code = self._extract_action_code(message)
@@ -333,21 +334,33 @@ class MessageService:
operator_user_id=login_user.user_id,
)
# 8. Auto-mark as read after action
await self.message_read_repository.mark_as_read(message_id, login_user.user_id)
# 8. Auto-mark as read after action (best-effort). The approval has already
# executed its handler and committed the status change above; the
# auto-read is a pure convenience side-step. A concurrent-insert race on
# the read record must never surface the finished approval as a 500.
try:
await self.message_read_repository.mark_as_read(message_id, login_user.user_id)
except IntegrityError:
logger.warning(
"Auto mark-as-read raced on approval message_id=%s user_id=%s; already read, ignoring",
message_id,
login_user.user_id,
)
logger.info(
"Approval action processed: message_id=%s, action=%s, operator=%s",
message_id, action.value, login_user.user_id,
message_id,
action.value,
login_user.user_id,
)
return updated_message
@staticmethod
def _update_content_after_approval(
content: List[Dict[str, Any]],
action: ApprovalActionEnum,
) -> List[Dict[str, Any]]:
content: list[dict[str, Any]],
action: ApprovalActionEnum,
) -> list[dict[str, Any]]:
"""
Update message content after approval action.
- Preserve 'user' and 'business_url' types for continued clickability
@@ -356,10 +369,10 @@ class MessageService:
updated = []
for item in content:
new_item = dict(item)
item_type = item.get('type', '')
item_type = item.get("type", "")
if item_type == 'agree_reject_button':
new_item['content'] = action.value
if item_type == "agree_reject_button":
new_item["content"] = action.value
updated.append(new_item)
@@ -378,19 +391,19 @@ class MessageService:
@staticmethod
def build_generic_notify_content(
content_list: List[MessageContentItem | Dict],
) -> List[Dict[str, Any]]:
content_list: list[MessageContentItem | dict],
) -> list[dict[str, Any]]:
"""
Build generic notification content.
"""
return [one if isinstance(one, dict) else one.to_message() for one in content_list]
async def send_generic_notify(
self,
sender: int,
receiver_user_ids: List[int],
content_item_list: List[MessageContentItem | Dict],
action_code: Optional[str] = None,
self,
sender: int,
receiver_user_ids: list[int],
content_item_list: list[MessageContentItem | dict],
action_code: str | None = None,
) -> InboxMessage:
"""
Send a generic notification message to specific receivers.
@@ -410,16 +423,16 @@ class MessageService:
@staticmethod
def build_generic_approval_content(
applicant_user_id: int,
applicant_user_name: str,
action_code: str,
business_type: str,
business_id: str,
business_name: str,
button_action_code: str,
approval_message_id: Optional[int] = None,
scenario_code: Optional[str] = None,
) -> List[Dict[str, Any]]:
applicant_user_id: int,
applicant_user_name: str,
action_code: str,
business_type: str,
business_id: str,
business_name: str,
button_action_code: str,
approval_message_id: int | None = None,
scenario_code: str | None = None,
) -> list[dict[str, Any]]:
"""
Build the generic message content structure for a business approval request.
"""
@@ -459,16 +472,16 @@ class MessageService:
return content
async def send_generic_approval(
self,
applicant_user_id: int,
applicant_user_name: str,
action_code: str,
business_type: str,
business_id: str,
business_name: str,
button_action_code: str,
receiver_user_ids: List[int],
scenario_code: Optional[str] = None,
self,
applicant_user_id: int,
applicant_user_name: str,
action_code: str,
business_type: str,
business_id: str,
business_name: str,
button_action_code: str,
receiver_user_ids: list[int],
scenario_code: str | None = None,
) -> InboxMessage:
"""
Send a generic approval notification to specific receivers.
@@ -499,12 +512,12 @@ class MessageService:
updated_content = []
for item in content:
new_item = dict(item)
if item.get('type') == 'agree_reject_button':
metadata = dict(item.get('metadata', {}))
data = dict(metadata.get('data', {}))
data['approval_id'] = str(message.id)
metadata['data'] = data
new_item['metadata'] = metadata
if item.get("type") == "agree_reject_button":
metadata = dict(item.get("metadata", {}))
data = dict(metadata.get("data", {}))
data["approval_id"] = str(message.id)
metadata["data"] = data
new_item["metadata"] = metadata
updated_content.append(new_item)
updated_message = await self.message_repository.update_message_content(message.id, updated_content)
@@ -522,21 +535,21 @@ class MessageService:
return message.action_code
# Fallback: extract from content JSON (supports both old key 'business_type' and new key 'action_code')
for item in (message.content or []):
if item.get('type') != 'agree_reject_button':
for item in message.content or []:
if item.get("type") != "agree_reject_button":
continue
metadata = item.get('metadata', {})
code = metadata.get('action_code') or metadata.get('business_type')
metadata = item.get("metadata", {})
code = metadata.get("action_code") or metadata.get("business_type")
if isinstance(code, str):
return code
return ""
async def batch_approve_channel_subscription_messages(
self,
channel_id: str,
operator_user_id: int,
self,
channel_id: str,
operator_user_id: int,
) -> int:
"""
Batch approve all pending channel subscription messages for a specific channel.
@@ -553,6 +566,8 @@ class MessageService:
if count > 0:
logger.info(
"Batch approved %d channel subscription messages for channel_id=%s, operator=%s",
count, channel_id, operator_user_id,
count,
channel_id,
operator_user_id,
)
return count
@@ -0,0 +1,152 @@
"""Concurrency regression tests for inbox message read-status writes.
Reproduces the production race condition (trace 19d54fcc… / 66d16bfb…):
``batch_mark_as_read`` / ``mark_as_read`` used a check-then-insert pattern with
no protection against a concurrent transaction inserting the same
``(message_id, user_id)`` between the SELECT and the INSERT, so the second
committer hit the ``ix_inbox_message_read_msg_user`` unique constraint
(dmPython IntegrityError -6602 / SQLAlchemy gkpj) and the whole batch rolled
back — every message in that request stayed unread and the API 500'd.
These tests drive two *independent* DB connections sharing one in-memory
database and inject a competing insert into the exact window between the
read-status SELECT and the batch INSERT, deterministically (no timing luck).
"""
from __future__ import annotations
import itertools
import pytest
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.pool import StaticPool
from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession
from bisheng.message.domain.models.inbox_message_read import InboxMessageRead
from bisheng.message.domain.repositories.implementations.inbox_message_read_repository_impl import (
InboxMessageReadRepositoryImpl,
)
# Each test gets a distinct shared-cache in-memory DB name so committed rows
# from one test never bleed into the next (a shared-cache DB lives as long as
# any connection to that name is open).
_db_counter = itertools.count()
def _make_engine(db_name: str):
# Shared-cache in-memory DB: multiple connections see the same tables/rows,
# but each connection runs its own transaction (unlike StaticPool ':memory:',
# which is a single connection with no isolation between sessions).
# StaticPool keeps exactly one live connection per engine so the shared
# in-memory DB is never torn down while a session is open.
uri = f"sqlite+aiosqlite:///file:{db_name}?mode=memory&cache=shared&uri=true"
return create_async_engine(
uri,
connect_args={"check_same_thread": False, "uri": True},
poolclass=StaticPool,
)
@pytest.fixture()
async def race_sessions():
"""Two AsyncSessions on independent connections over one shared in-memory DB."""
db_name = f"inbox_read_race_{next(_db_counter)}"
engine_a = _make_engine(db_name)
engine_b = _make_engine(db_name)
async with engine_a.begin() as conn:
await conn.run_sync(InboxMessageRead.__table__.create)
session_a = AsyncSession(bind=engine_a, expire_on_commit=False)
session_b = AsyncSession(bind=engine_b, expire_on_commit=False)
try:
yield session_a, session_b
finally:
await session_a.close()
await session_b.close()
await engine_a.dispose()
await engine_b.dispose()
async def _count_rows(session: AsyncSession, message_id: int, user_id: int) -> int:
result = await session.exec(
select(InboxMessageRead).where(
InboxMessageRead.message_id == message_id,
InboxMessageRead.user_id == user_id,
)
)
return len(list(result.all()))
def _inject_competing_insert_after_first_query(session_a, session_b, message_id, user_id):
"""Patch session_a.exec so that, right after the FIRST query returns (the
read-status SELECT), a competing transaction on session_b inserts
``(message_id, user_id)`` and commits — landing us in the exact race window.
Only the first exec is hijacked; the repository's own retry re-query runs
unhijacked so it can observe the competitor's row.
"""
state = {"fired": False}
orig_exec = session_a.exec
async def exec_with_injection(*args, **kwargs):
result = await orig_exec(*args, **kwargs)
if not state["fired"]:
state["fired"] = True
session_b.add(InboxMessageRead(message_id=message_id, user_id=user_id))
await session_b.commit()
return result
session_a.exec = exec_with_injection
async def test_batch_mark_as_read_is_idempotent_under_concurrent_insert(race_sessions):
session_a, session_b = race_sessions
user_id = 1
shared_id = 123 # both requests try to mark this one
only_a_id = 456 # only request A marks this one
repo_a = InboxMessageReadRepositoryImpl(session_a)
_inject_competing_insert_after_first_query(session_a, session_b, shared_id, user_id)
# Must NOT raise IntegrityError even though `shared_id` gets inserted by the
# competitor in the SELECT→INSERT window.
newly_marked = await repo_a.batch_mark_as_read([shared_id, only_a_id], user_id)
# A only newly marked `only_a_id`; `shared_id` was won by the competitor.
assert newly_marked == 1
# Exactly one row per (message_id, user_id) — no duplicates, nothing lost.
assert await _count_rows(session_a, shared_id, user_id) == 1
assert await _count_rows(session_a, only_a_id, user_id) == 1
async def test_mark_as_read_is_idempotent_under_concurrent_insert(race_sessions):
session_a, session_b = race_sessions
user_id = 7
shared_id = 999
repo_a = InboxMessageReadRepositoryImpl(session_a)
_inject_competing_insert_after_first_query(session_a, session_b, shared_id, user_id)
# find_one (the first exec) returns empty, competitor inserts, then save()
# would hit the unique constraint. Must be swallowed idempotently.
record = await repo_a.mark_as_read(shared_id, user_id)
assert record is not None
assert record.message_id == shared_id
assert record.user_id == user_id
assert await _count_rows(session_a, shared_id, user_id) == 1
async def test_batch_mark_as_read_dedupes_duplicate_ids(race_sessions):
"""Duplicate ids within one request must not self-collide on the unique index."""
session_a, _ = race_sessions
user_id = 3
repo_a = InboxMessageReadRepositoryImpl(session_a)
newly_marked = await repo_a.batch_mark_as_read([555, 555, 666], user_id)
assert newly_marked == 2
assert await _count_rows(session_a, 555, user_id) == 1
assert await _count_rows(session_a, 666, user_id) == 1