Files
banban/mini-program/app/service/im.py

708 lines
22 KiB
Python

from dataclasses import dataclass
import json
from collections.abc import Mapping
from typing import Any
from fastapi import HTTPException, status
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
try:
from app.db_compat import current_timestamp_sql, inserted_primary_key, select_for_update_clause
from app.schemas.im import (
ChildConversationMessageItem,
DeviceMessageCreateRequest,
ParentChildMessageCreateRequest,
)
except ModuleNotFoundError:
from db_compat import current_timestamp_sql, inserted_primary_key, select_for_update_clause
from schemas.im import ChildConversationMessageItem, DeviceMessageCreateRequest, ParentChildMessageCreateRequest
PARENT_PARTICIPANT_TYPE = 1
CHILD_PARTICIPANT_TYPE = 2
CHILD_PEER_CONVERSATION_TYPE = 1
PARENT_CHILD_CONVERSATION_TYPE = 2
CONVERSATION_TYPE_NAMES = {
CHILD_PEER_CONVERSATION_TYPE: "child_peer",
PARENT_CHILD_CONVERSATION_TYPE: "parent_child",
}
PARTICIPANT_TYPE_NAMES = {
PARENT_PARTICIPANT_TYPE: "parent",
CHILD_PARTICIPANT_TYPE: "child",
}
@dataclass(frozen=True)
class DeviceIdentity:
device_id: str
child_id: int
child_name: str | None
@dataclass(frozen=True)
class ConversationMessageCreateResult:
idempotent: bool
conversation_id: int
conversation_type: int
message: ChildConversationMessageItem
@property
def conversation_type_name(self) -> str:
return conversation_type_name(self.conversation_type)
def conversation_type_name(conversation_type: int) -> str:
return CONVERSATION_TYPE_NAMES.get(conversation_type, f"unknown_{conversation_type}")
def participant_type_name(participant_type: int) -> str:
return PARTICIPANT_TYPE_NAMES.get(participant_type, f"unknown_{participant_type}")
def normalize_content_json(value: Any) -> dict[str, Any] | None:
if value is None:
return None
if isinstance(value, dict):
return value
if isinstance(value, str):
try:
parsed = json.loads(value)
except json.JSONDecodeError:
return None
if isinstance(parsed, dict):
return parsed
return None
def row_to_message_item(row: Mapping[str, Any]) -> ChildConversationMessageItem:
return ChildConversationMessageItem(
id=int(row["id"]),
conversation_id=int(row["conversation_id"]),
seq=int(row["seq"]),
sender_type=participant_type_name(int(row["sender_type"])),
sender_id=str(row["sender_id"]),
receiver_type=participant_type_name(int(row["receiver_type"])),
receiver_id=str(row["receiver_id"]),
content_type=int(row["content_type"]),
content_text=row["content_text"],
content_json=normalize_content_json(row["content_json"]),
media_file_key=row["media_file_key"],
media_duration_ms=row["media_duration_ms"],
media_mime_type=row["media_mime_type"],
media_size_bytes=row["media_size_bytes"],
media_transcript_text=row["media_transcript_text"],
client_msg_id=row["client_msg_id"],
sender_name_snapshot=row["sender_name_snapshot"],
sender_avatar_snapshot=row["sender_avatar_snapshot"],
receiver_name_snapshot=row["receiver_name_snapshot"],
receiver_avatar_snapshot=row["receiver_avatar_snapshot"],
ext_json=normalize_content_json(row["ext_json"]),
created_at=row["created_at"],
)
def assert_parent_child_access(db: Session, *, user_id: int, child_id: int) -> Mapping[str, Any]:
child_row = _get_child_row(db, child_id)
if not child_row:
raise HTTPException(status_code=404, detail="child not found")
has_access = (
db.execute(
text(
"""
SELECT 1
FROM parent_child_relations
WHERE user_id = :user_id
AND child_id = :child_id
AND status = 1
LIMIT 1
"""
),
{"user_id": user_id, "child_id": child_id},
).scalar_one_or_none()
is not None
)
if not has_access:
raise HTTPException(status_code=403, detail="no permission to access this child")
return child_row
def authenticate_device_identity(db: Session, *, device_id: str, serial_number: str) -> DeviceIdentity:
row = (
db.execute(
text(
"""
SELECT
da.device_id,
db.child_id,
c.child_name
FROM device_auth AS da
LEFT JOIN device_bindings AS db
ON db.device_id = da.device_id
AND db.status = 1
LEFT JOIN children AS c
ON c.child_id = db.child_id
AND c.status = 1
WHERE da.device_id = :device_id
AND da.serial_number = :serial_number
AND da.is_active = 1
LIMIT 1
"""
),
{"device_id": device_id, "serial_number": serial_number},
)
.mappings()
.first()
)
if not row:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid device credentials")
if row["child_id"] is None:
raise HTTPException(status_code=404, detail="device not bound to a child")
return DeviceIdentity(
device_id=str(row["device_id"]),
child_id=int(row["child_id"]),
child_name=row["child_name"],
)
def create_parent_child_message(
db: Session,
*,
parent_user_id: int,
child_id: int,
payload: ParentChildMessageCreateRequest,
) -> ConversationMessageCreateResult:
child_row = assert_parent_child_access(db, user_id=parent_user_id, child_id=child_id)
parent_row = _get_parent_row(db, parent_user_id)
if not parent_row:
raise HTTPException(status_code=404, detail="parent not found")
conversation_id, idempotent = _create_message(
db=db,
conversation_type=PARENT_CHILD_CONVERSATION_TYPE,
participant_a_type=CHILD_PARTICIPANT_TYPE,
participant_a_id=str(child_id),
participant_b_type=PARENT_PARTICIPANT_TYPE,
participant_b_id=str(parent_user_id),
pair_key=f"{child_id}:{parent_user_id}",
sender_type=PARENT_PARTICIPANT_TYPE,
sender_id=str(parent_user_id),
receiver_type=CHILD_PARTICIPANT_TYPE,
receiver_id=str(child_id),
sender_name_snapshot=parent_row["nickname"],
sender_avatar_snapshot=parent_row["avatar_url"],
receiver_name_snapshot=child_row["child_name"],
receiver_avatar_snapshot=None,
payload=payload,
)
message_row = _get_message_by_conversation_client_id(
db,
conversation_id=conversation_id,
client_msg_id=payload.client_msg_id,
)
if not message_row:
raise RuntimeError("message was not found after insert")
return ConversationMessageCreateResult(
idempotent=idempotent,
conversation_id=conversation_id,
conversation_type=PARENT_CHILD_CONVERSATION_TYPE,
message=row_to_message_item(message_row),
)
def create_device_message(
db: Session,
*,
device_id: str,
serial_number: str,
payload: DeviceMessageCreateRequest,
) -> tuple[DeviceIdentity, ConversationMessageCreateResult]:
device_identity = authenticate_device_identity(
db,
device_id=device_id,
serial_number=serial_number,
)
if payload.conversation_type == CHILD_PEER_CONVERSATION_TYPE:
if payload.peer_child_id == device_identity.child_id:
raise HTTPException(status_code=400, detail="peer_child_id must be different from current child")
sender_child_row = assert_child_exists(db, child_id=device_identity.child_id)
receiver_child_row = assert_child_exists(db, child_id=payload.peer_child_id)
participant_a_id, participant_b_id, pair_key = _build_child_peer_pair(
device_identity.child_id,
payload.peer_child_id,
)
conversation_id, idempotent = _create_message(
db=db,
conversation_type=CHILD_PEER_CONVERSATION_TYPE,
participant_a_type=CHILD_PARTICIPANT_TYPE,
participant_a_id=participant_a_id,
participant_b_type=CHILD_PARTICIPANT_TYPE,
participant_b_id=participant_b_id,
pair_key=pair_key,
sender_type=CHILD_PARTICIPANT_TYPE,
sender_id=str(device_identity.child_id),
receiver_type=CHILD_PARTICIPANT_TYPE,
receiver_id=str(payload.peer_child_id),
sender_name_snapshot=sender_child_row["child_name"],
sender_avatar_snapshot=None,
receiver_name_snapshot=receiver_child_row["child_name"],
receiver_avatar_snapshot=None,
payload=payload,
)
else:
sender_child_row = assert_child_exists(db, child_id=device_identity.child_id)
parent_row = _get_parent_row(db, payload.parent_user_id)
if not parent_row:
raise HTTPException(status_code=404, detail="parent not found")
assert_parent_child_access(db, user_id=payload.parent_user_id, child_id=device_identity.child_id)
conversation_id, idempotent = _create_message(
db=db,
conversation_type=PARENT_CHILD_CONVERSATION_TYPE,
participant_a_type=CHILD_PARTICIPANT_TYPE,
participant_a_id=str(device_identity.child_id),
participant_b_type=PARENT_PARTICIPANT_TYPE,
participant_b_id=str(payload.parent_user_id),
pair_key=f"{device_identity.child_id}:{payload.parent_user_id}",
sender_type=CHILD_PARTICIPANT_TYPE,
sender_id=str(device_identity.child_id),
receiver_type=PARENT_PARTICIPANT_TYPE,
receiver_id=str(payload.parent_user_id),
sender_name_snapshot=sender_child_row["child_name"],
sender_avatar_snapshot=None,
receiver_name_snapshot=parent_row["nickname"],
receiver_avatar_snapshot=parent_row["avatar_url"],
payload=payload,
)
message_row = _get_message_by_conversation_client_id(
db,
conversation_id=conversation_id,
client_msg_id=payload.client_msg_id,
)
if not message_row:
raise RuntimeError("message was not found after insert")
result = ConversationMessageCreateResult(
idempotent=idempotent,
conversation_id=conversation_id,
conversation_type=payload.conversation_type,
message=row_to_message_item(message_row),
)
return device_identity, result
def assert_child_exists(db: Session, *, child_id: int) -> Mapping[str, Any]:
child_row = _get_child_row(db, child_id)
if not child_row:
raise HTTPException(status_code=404, detail="child not found")
return child_row
def _get_child_row(db: Session, child_id: int) -> Mapping[str, Any] | None:
return (
db.execute(
text(
"""
SELECT child_id, child_name, status
FROM children
WHERE child_id = :child_id
AND status = 1
LIMIT 1
"""
),
{"child_id": child_id},
)
.mappings()
.first()
)
def _get_parent_row(db: Session, user_id: int) -> Mapping[str, Any] | None:
return (
db.execute(
text(
"""
SELECT user_id, nickname, avatar_url, status
FROM parents
WHERE user_id = :user_id
AND status = 1
LIMIT 1
"""
),
{"user_id": user_id},
)
.mappings()
.first()
)
def _build_child_peer_pair(child_a_id: int, child_b_id: int) -> tuple[str, str, str]:
low_id, high_id = sorted((child_a_id, child_b_id))
participant_a_id = str(low_id)
participant_b_id = str(high_id)
return participant_a_id, participant_b_id, f"{participant_a_id}:{participant_b_id}"
def _build_preview(content_type: int, content_text: str | None) -> str:
if content_type == 1:
return (content_text or "").strip()[:255]
if content_type == 2:
return "[audio]"
if content_type == 3:
return "[image]"
return "[json]"
def _create_message(
db: Session,
*,
conversation_type: int,
participant_a_type: int,
participant_a_id: str,
participant_b_type: int,
participant_b_id: str,
pair_key: str,
sender_type: int,
sender_id: str,
receiver_type: int,
receiver_id: str,
sender_name_snapshot: str | None,
sender_avatar_snapshot: str | None,
receiver_name_snapshot: str | None,
receiver_avatar_snapshot: str | None,
payload: ParentChildMessageCreateRequest | DeviceMessageCreateRequest,
) -> tuple[int, bool]:
conversation_id = _get_or_create_conversation(
db=db,
conversation_type=conversation_type,
participant_a_type=participant_a_type,
participant_a_id=participant_a_id,
participant_b_type=participant_b_type,
participant_b_id=participant_b_id,
pair_key=pair_key,
)
existing = _get_message_by_conversation_client_id(
db,
conversation_id=conversation_id,
client_msg_id=payload.client_msg_id,
)
if existing:
return conversation_id, True
now_sql = current_timestamp_sql(db)
preview = _build_preview(payload.content_type, payload.content_text)
try:
conversation_row = _get_conversation_by_id(db, conversation_id=conversation_id, lock=True)
if not conversation_row:
raise HTTPException(status_code=404, detail="conversation not found")
next_seq = int(conversation_row["last_seq"]) + 1
message_id = _next_primary_key(db, "im_messages")
insert_sql = f"""
INSERT INTO im_messages (
{'id,' if message_id is not None else ''}
conversation_id,
seq,
sender_type,
sender_id,
receiver_type,
receiver_id,
content_type,
content_text,
content_json,
media_file_key,
media_duration_ms,
media_mime_type,
media_size_bytes,
media_transcript_text,
client_msg_id,
sender_name_snapshot,
sender_avatar_snapshot,
receiver_name_snapshot,
receiver_avatar_snapshot,
ext_json,
created_at
) VALUES (
{':id,' if message_id is not None else ''}
:conversation_id,
:seq,
:sender_type,
:sender_id,
:receiver_type,
:receiver_id,
:content_type,
:content_text,
:content_json,
:media_file_key,
:media_duration_ms,
:media_mime_type,
:media_size_bytes,
:media_transcript_text,
:client_msg_id,
:sender_name_snapshot,
:sender_avatar_snapshot,
:receiver_name_snapshot,
:receiver_avatar_snapshot,
:ext_json,
{now_sql}
)
"""
params: dict[str, Any] = {
"conversation_id": conversation_id,
"seq": next_seq,
"sender_type": sender_type,
"sender_id": sender_id,
"receiver_type": receiver_type,
"receiver_id": receiver_id,
"content_type": payload.content_type,
"content_text": payload.content_text,
"content_json": json.dumps(payload.content_json) if payload.content_json is not None else None,
"media_file_key": payload.media_file_key,
"media_duration_ms": payload.media_duration_ms,
"media_mime_type": payload.media_mime_type,
"media_size_bytes": payload.media_size_bytes,
"media_transcript_text": payload.media_transcript_text,
"client_msg_id": payload.client_msg_id,
"sender_name_snapshot": sender_name_snapshot,
"sender_avatar_snapshot": sender_avatar_snapshot,
"receiver_name_snapshot": receiver_name_snapshot,
"receiver_avatar_snapshot": receiver_avatar_snapshot,
"ext_json": json.dumps(payload.ext_json) if payload.ext_json is not None else None,
}
if message_id is not None:
params["id"] = message_id
result = db.execute(text(insert_sql), params)
if message_id is None:
message_id = inserted_primary_key(result)
db.execute(
text(
f"""
UPDATE im_conversations
SET status = 1,
last_seq = :last_seq,
message_count = message_count + 1,
last_message_preview = :last_message_preview,
last_message_at = {now_sql},
updated_at = {now_sql}
WHERE id = :conversation_id
"""
),
{
"conversation_id": conversation_id,
"last_seq": next_seq,
"last_message_preview": preview,
},
)
db.commit()
return conversation_id, False
except IntegrityError:
db.rollback()
existing = _get_message_by_conversation_client_id(
db,
conversation_id=conversation_id,
client_msg_id=payload.client_msg_id,
)
if existing:
return conversation_id, True
raise
except Exception:
db.rollback()
raise
def _get_or_create_conversation(
db: Session,
*,
conversation_type: int,
participant_a_type: int,
participant_a_id: str,
participant_b_type: int,
participant_b_id: str,
pair_key: str,
) -> int:
row = _get_conversation_by_pair(
db,
conversation_type=conversation_type,
pair_key=pair_key,
lock=False,
)
if row:
return int(row["id"])
now_sql = current_timestamp_sql(db)
conversation_id = _next_primary_key(db, "im_conversations")
insert_sql = f"""
INSERT INTO im_conversations (
{'id,' if conversation_id is not None else ''}
conversation_type,
participant_a_type,
participant_a_id,
participant_b_type,
participant_b_id,
pair_key,
status,
last_seq,
message_count,
created_at,
updated_at
) VALUES (
{':id,' if conversation_id is not None else ''}
:conversation_type,
:participant_a_type,
:participant_a_id,
:participant_b_type,
:participant_b_id,
:pair_key,
1,
0,
0,
{now_sql},
{now_sql}
)
"""
params = {
"conversation_type": conversation_type,
"participant_a_type": participant_a_type,
"participant_a_id": participant_a_id,
"participant_b_type": participant_b_type,
"participant_b_id": participant_b_id,
"pair_key": pair_key,
}
if conversation_id is not None:
params["id"] = conversation_id
try:
result = db.execute(text(insert_sql), params)
if conversation_id is not None:
return conversation_id
return inserted_primary_key(result)
except IntegrityError:
db.rollback()
row = _get_conversation_by_pair(
db,
conversation_type=conversation_type,
pair_key=pair_key,
lock=False,
)
if row:
return int(row["id"])
raise
def _get_conversation_by_pair(
db: Session,
*,
conversation_type: int,
pair_key: str,
lock: bool,
) -> Mapping[str, Any] | None:
lock_clause = select_for_update_clause(db) if lock else ""
return (
db.execute(
text(
f"""
SELECT id, conversation_type, last_seq, status
FROM im_conversations
WHERE conversation_type = :conversation_type
AND pair_key = :pair_key
LIMIT 1{lock_clause}
"""
),
{"conversation_type": conversation_type, "pair_key": pair_key},
)
.mappings()
.first()
)
def _get_conversation_by_id(
db: Session,
*,
conversation_id: int,
lock: bool,
) -> Mapping[str, Any] | None:
lock_clause = select_for_update_clause(db) if lock else ""
return (
db.execute(
text(
f"""
SELECT id, last_seq, message_count, status
FROM im_conversations
WHERE id = :conversation_id
LIMIT 1{lock_clause}
"""
),
{"conversation_id": conversation_id},
)
.mappings()
.first()
)
def _get_message_by_conversation_client_id(
db: Session,
*,
conversation_id: int,
client_msg_id: str,
) -> Mapping[str, Any] | None:
return (
db.execute(
text(
"""
SELECT
id,
conversation_id,
seq,
sender_type,
sender_id,
receiver_type,
receiver_id,
content_type,
content_text,
content_json,
media_file_key,
media_duration_ms,
media_mime_type,
media_size_bytes,
media_transcript_text,
client_msg_id,
sender_name_snapshot,
sender_avatar_snapshot,
receiver_name_snapshot,
receiver_avatar_snapshot,
ext_json,
created_at
FROM im_messages
WHERE conversation_id = :conversation_id
AND client_msg_id = :client_msg_id
LIMIT 1
"""
),
{"conversation_id": conversation_id, "client_msg_id": client_msg_id},
)
.mappings()
.first()
)
def _next_primary_key(db: Session, table_name: str) -> int | None:
if db.get_bind() is None or db.get_bind().dialect.name != "sqlite":
return None
return int(
db.execute(text(f"SELECT COALESCE(MAX(id), 0) + 1 FROM {table_name}")).scalar_one()
)